Diff Coverage

Diff: origin/master...HEAD, staged and unstaged changes

Source File Diff Coverage (%) Missing Lines
hyper_parallel/compile/__init__.py 100%  
hyper_parallel/compile/dependency_bridge.py 0.0% 17-18,20-21,23,44,52,60,72-74,77,79,86,98-100,104-105,115,117-118,122
hyper_parallel/compile/text_trainer.py 8.6% 27,29,34-36,39,42,51-52,55-58,67,70,72,79-84,86,88-89,91-92,94-95,97-100,102,105,107,110,112,114-115,122-123,125-127,129,135-137,139,145-146,148-151,153-155,157-158,166-167,172-173,179,181,187-189,193-195,198
hyper_parallel/compile/tracer/dynamic_shapes.py 29.7% 53-55,57,61-62,64-66,68-72,74,82,97,99-103,105-107,111
hyper_parallel/compile/tracer/graph_tracer.py 50.0% 602,605
hyper_parallel/compile/trainer.py 30.0% 83,128-129,217-218,326,328
hyper_parallel/trainer/config/__init__.py 100%  
hyper_parallel/trainer/config/graph.py 100%  
hyper_parallel/trainer/config/trainer.py 100%  
hyper_parallel/compile/dependency_bridge.py
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
# limitations under the License.
# ============================================================================
"""Stable bridge helpers from graph mode to the eager Trainer/AutoModel stack."""

from dataclasses import replace
from typing import Any, Optional

from hyper_parallel.models.build_options import CompileConfig
from hyper_parallel.trainer.config import Target, TrainerConfig

from .parallel_config import PassConfig


# def _wrap_dataloader_for_graph_mode(dataloader_config: Any) -> Any:
#     """Enable graph-stable padding for dynamic online text batches."""
40
41
42
43
44
45
46
47
48
#         target=target.replace(pad_to_token_budget=True),
#     )


def clone_config_for_graph_mode(config: TrainerConfig) -> TrainerConfig:
    """Clone a TrainerConfig and disable eager per-layer compile.

    Phase-1 graph integration reuses AutoModel for model preparation, but the
    execution step belongs to ``hyper_parallel.compile``. The eager
48
49
50
51
52
53
54
55
56
    execution step belongs to ``hyper_parallel.compile``. The eager
    ``distributed.compile.apply_compile`` path must therefore be disabled to
    avoid compiling decoder layers twice.
    """
    return replace(
        config,
        model=wrap_model_target_for_graph_mode(config.model),
        # dataloader=_wrap_dataloader_for_graph_mode(config.dataloader),
        compile=CompileConfig(enabled=False),
56
57
58
59
60
61
62
63
64
        compile=CompileConfig(enabled=False),
    )


def build_model_for_graph_mode(
    *,
    model_target: Target[Any],
    distributed_setup: Any = None,
    **kwargs: Any,
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
    Graph mode still wants AutoModel's mesh construction, TP/SP sharding plan,
    and other preparation steps. The eager FSDP2 runtime wrap, however, must be
    disabled so the graph pass pipeline owns FSDP behavior.
    """
    if distributed_setup is not None:
        distributed_setup = replace(distributed_setup, strategy_config=None)
    return model_target.build(distributed_setup=distributed_setup, **kwargs)


def wrap_model_target_for_graph_mode(model_target: Target[Any]) -> Target[Any]:
    """Wrap the configured model target with graph-mode runtime adjustments."""
    return Target(
        _target_=build_model_for_graph_mode,
        target_path="hyper_parallel.compile.dependency_bridge.build_model_for_graph_mode",
        model_target=model_target,
    )
82
83
84
85
86
87
88
89
90
        model_target=model_target,
    )


def build_pass_config_from_trainer_config(
    config: TrainerConfig,
    *,
    fsdp_enabled: Optional[bool] = None,
    enable_overlap: bool = False,
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
    The first integration milestone keeps graph mode focused on model
    preparation reuse plus single-process execution, so FSDP graph passes stay
    opt-in and disabled by default.
    """
    accelerator = config.accelerator
    if fsdp_enabled is None:
        fsdp_enabled = (
            config.fsdp_config.dp_shard_size > 1
            or config.fsdp_config.edp_shard_size > 1
        )
    fsdp_degree = config.fsdp_config.dp_shard_size if fsdp_enabled else None
    return PassConfig(
        enable_overlap=enable_overlap,
        fsdp_enabled=fsdp_enabled,
        fsdp_degree=fsdp_degree,
        tp_size=accelerator.tp_size,
111
112
113
114
115
116
117
118
119
120
121
122
        loss_parallel=accelerator.loss_parallel,
    )


def loss_to_metrics(loss: Any) -> dict[str, Any]:
    """Normalize graph loss output to a callback-friendly metrics mapping."""
    if isinstance(loss, dict):
        return {
            str(name): value.detach() if hasattr(value, "detach") else value
            for name, value in loss.items()
        }
    return {"graph_loss": loss.detach() if hasattr(loss, "detach") else loss}
hyper_parallel/compile/text_trainer.py
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
from torch._dynamo import mark_dynamic

from hyper_parallel.trainer.runtime.loss_aggregation import count_loss_token
from hyper_parallel.trainer.text_trainer import TextTrainer
from hyper_parallel.trainer.config import TrainerConfig

from .dependency_bridge import (
    build_pass_config_from_trainer_config,
    clone_config_for_graph_mode,
    loss_to_metrics,
)
from .parallel_config import PassConfig
from .sharding_config import PassPlan
from .trainer import GraphTrainer as GraphExecutionEngine


class GraphTextTrainer(TextTrainer):
    """Reuse TextTrainer and replace only the forward/backward step with graph mode."""

    def __init__(
        self,
        config: TrainerConfig,
        *,
        train_fn: Optional[Any] = None,
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
        pass_config: Optional[PassConfig] = None,
        pass_plan: Optional[PassPlan] = None,
    ) -> None:
        """Build graph-mode text training on top of the eager TextTrainer stages."""
        graph_config = clone_config_for_graph_mode(config)
        self.graph_pass_config = pass_config or build_pass_config_from_trainer_config(
            graph_config
        )
        self.graph_pass_plan = pass_plan
        self._graph_train_fn = train_fn or self._default_train_fn
        super().__init__(graph_config)
        self.graph_executor = GraphExecutionEngine(
            model=self.base.model,
            train_fn=self._graph_train_fn,
            pass_config=self.graph_pass_config,
            pass_plan=self.graph_pass_plan,
63
64
65
66
67
68
69
70
71
72
73
74
75
76
            device=self.base.device,
            mesh_context=self.base.mesh,
            manage_optimizer=False,
        )
        self._dynamic_shape_annotation_enabled = self._should_enable_dynamic_shapes(
            graph_config
        )
        self._dynamic_shapes_marked = False

    def _default_train_fn(
        self,
        model: torch.nn.Module,
        model_inputs: Mapping[str, Any],
        loss_inputs: Mapping[str, Any],
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
        model_inputs: Mapping[str, Any],
        loss_inputs: Mapping[str, Any],
    ) -> torch.Tensor:
        """Default graph trace function for Transformer-style text training."""
        outputs = model(**dict(model_inputs), use_cache=False)
        labels = loss_inputs.get("labels")
        loss = self.base.loss_fn(model_output=outputs, labels=labels)
        if isinstance(loss, dict):
            return torch.stack(list(loss.values())).sum()
        return loss

    def set_pytree_pre_hook(self, hook: Any) -> "GraphTextTrainer":
        """Register a tracer pre-hook on the underlying graph executor."""
        self.graph_executor.set_pytree_pre_hook(hook)
        return self

    @staticmethod
    def _should_enable_dynamic_shapes(config: TrainerConfig) -> bool:
        """Return whether the limited online-text dynamic-shape path applies."""
        if not bool(getattr(config.graph, "enabled", False)):
            return False

        dataloader_target = getattr(config.dataloader, "target", None)
        get_batch_target = getattr(config.dataloader, "get_batch", None)
        if dataloader_target is None or get_batch_target is None:
            return False

        if getattr(dataloader_target, "_target_path", None) != (
            "hyper_parallel.data.batching.DynamicBatchDataLoader"
        ):
            return False

        if getattr(get_batch_target, "_target_path", None) != (
            "hyper_parallel.data.batching.ParallelBatch"
        ):
            return False

        return getattr(get_batch_target, "source_type", None) == "online"

    @staticmethod
    def _mark_tensor_token_dim(
        tensor: Any,
        *,
        token_budget: int,
        shape_id: str,
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
        token_budget: int,
        shape_id: str,
    ) -> None:
        """Mark the token dimension of one plain Tensor input as dynamic."""
        if not isinstance(tensor, torch.Tensor) or tensor.dim() == 0:
            return

        token_dim = tensor.dim() - 1
        if tensor.shape[token_dim] <= 0:
            return

        mark_dynamic(
            tensor,
            token_dim,
            min=1,
            max=token_budget,
131
132
133
134
135
136
137
138
139
140
141
142
143
            token_dim,
            min=1,
            max=token_budget,
        )
        shape_ids = dict(getattr(tensor, "_dynamo_shape_ids", {}) or {})
        shape_ids[token_dim] = shape_id
        setattr(tensor, "_dynamo_shape_ids", shape_ids)

    def _mark_online_text_dynamic_shapes(
        self,
        model_inputs: Mapping[str, Any],
        loss_inputs: Mapping[str, Any],
    ) -> None:
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
        model_inputs: Mapping[str, Any],
        loss_inputs: Mapping[str, Any],
    ) -> None:
        """Mark the packed token dimension for the first graph trace batch."""
        if not self._dynamic_shape_annotation_enabled or self._dynamic_shapes_marked:
            return

        max_seq_len = getattr(self.base.config.dataset.data_transform, "max_seq_len", None)
        micro_batch_size = getattr(self.base.config.training, "micro_batch_size", None)
        if max_seq_len is None or micro_batch_size is None:
            return

        token_budget = int(max_seq_len) * int(micro_batch_size)
        if token_budget <= 0:
            return

        shape_id = "graph_text_tokens"
        for field in (
            "input_ids",
            "labels",
            "position_ids",
            "attention_mask",
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
            "attention_mask",
            "shift_labels",
            "loss_mask",
        ):
            if field in model_inputs:
                self._mark_tensor_token_dim(
                    model_inputs[field],
                    token_budget=token_budget,
                    shape_id=shape_id,
                )
            if field in loss_inputs:
                self._mark_tensor_token_dim(
                    loss_inputs[field],
                    token_budget=token_budget,
                    shape_id=shape_id,
                )
175
176
177
178
179
180
181
182
183
184
185
                    token_budget=token_budget,
                    shape_id=shape_id,
                )

        self._dynamic_shapes_marked = True

    def forward_backward_step(
        self,
        data_iterator: Any,
        num_micro_steps: int,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
        data_iterator: Any,
        num_micro_steps: int,
    ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
        """Fetch one text batch and execute graph-mode forward/backward."""
        model_inputs, loss_inputs = self.base.get_batch(data_iterator)
        self.base.current_token_counts = count_loss_token(loss_inputs)
        self.base.step_token_counts = {
            name: token_count * num_micro_steps
            for name, token_count in self.base.current_token_counts.items()
        }
        self._mark_online_text_dynamic_shapes(model_inputs, loss_inputs)
        loss = self.graph_executor.train_step(model_inputs, loss_inputs)
        return loss, loss_to_metrics(loss)


__all__ = ["GraphTextTrainer"]
hyper_parallel/compile/tracer/dynamic_shapes.py
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78


def _symbolic_context_for_marked_dims(tensor: torch.Tensor) -> Any | None:
    """Translate ``mark_dynamic`` metadata into a ``StatelessSymbolicContext``."""
    marked_dynamic_indices = getattr(tensor, "_dynamo_dynamic_indices", set())
    if not marked_dynamic_indices:
        return None

    dynamic_ranges = {
        dim_range.dim: dim_range
        for dim_range in getattr(tensor, "_dynamo_dynamic_range", set())
    }
    dynamic_sizes = [DimDynamic.STATIC] * tensor.dim()
    constraint_sizes: list[Any] = [None] * tensor.dim()

    for dim in range(tensor.dim()):
        if dim not in marked_dynamic_indices:
            continue

        dynamic_sizes[dim] = DimDynamic.DYNAMIC
        dim_range = dynamic_ranges.get(dim)
        if dim_range is None or (dim_range.min is None and dim_range.max is None):
            constraint_sizes[dim] = RelaxedUnspecConstraint(warn_only=False)
            continue

        constraint_sizes[dim] = StrictMinMaxConstraint(
            vr=ValueRanges(
                lower=0 if dim_range.min is None else dim_range.min,
                upper=int_oo if dim_range.max is None else dim_range.max,
            ),
78
79
80
81
82
83
84
85
86
            ),
            warn_only=False,
        )

    return StatelessSymbolicContext(
        dynamic_sizes=dynamic_sizes,
        constraint_sizes=constraint_sizes,
        shape_ids=getattr(tensor, "_dynamo_shape_ids", None),
    )
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
    input_name: str,
) -> torch.Tensor:
    """Fakeify one tracer input while preserving dynamic-shape annotations."""
    # pylint: disable=C0415
    from torch._dynamo.source import LocalSource

    def copy_tensor_annotations(fake_tensor: torch.Tensor) -> torch.Tensor:
        for name in _DYNAMO_SHAPE_ANNOTATION_NAMES:
            if hasattr(tensor, name):
                setattr(fake_tensor, name, getattr(tensor, name))
        return fake_tensor

    symbolic_context = _symbolic_context_for_marked_dims(tensor)
    if symbolic_context is None:
        return copy_tensor_annotations(
            fake_mode.from_tensor(tensor, static_shapes=True)
        )

    return copy_tensor_annotations(
        fake_mode.from_tensor(
            tensor,
            source=LocalSource(input_name, is_input=True),
            symbolic_context=symbolic_context,
hyper_parallel/compile/tracer/graph_tracer.py
598
599
600
601
602
603
604
605
606
607
608
        )

    state_flat, _ = torch.utils._pytree.tree_flatten({"model": model_state})

    user_inputs_flat, _ = torch.utils._pytree.tree_flatten(
        (input_batch, label_batch)
    )
    flat_inputs = list(state_flat) + list(user_inputs_flat)

    with torch.no_grad():
        outputs = joint_graph.graph_module(*flat_inputs)
hyper_parallel/compile/trainer.py
79
80
81
82
83
84
85
86
87
        self.pass_config = pass_config
        self.pass_plan = pass_plan
        self.optimizer_config = optimizer_config or {}
        self._mesh_context = mesh_context
        self._manage_optimizer = manage_optimizer
        self.device = device or (
            torch.device("npu")
            if (hasattr(torch, "npu") and torch.npu.is_available())
            else torch.device("cpu")
124
125
126
127
128
129
130
131
132
        pipeline.run(joint_graph.graph_module, **pass_kwargs)

        self._joint_graph = joint_graph

        if self._manage_optimizer:
            self._init_optimizer()

    def _init_device_mesh(self, mesh_context: Optional[Any] = None):
        """Initialize the FSDP process group.
213
214
215
216
217
218
219
220
221
222
        return loss

    def optimizer_step(self) -> None:
        """Optimizer update"""
        if not self._manage_optimizer:
            raise RuntimeError(
                "optimizer_step() is disabled when manage_optimizer=False. "
                "Let the outer trainer runtime own optimizer stepping."
            )
        if self.optimizer is None:
322
323
324
325
326
327
328
329
330
331
        model preparation still flows through AutoModel + Trainer config,
        while the forward/backward step is executed by the graph tracer and
        pass pipeline.
        """
        from .text_trainer import GraphTextTrainer  # pylint: disable=C0415

        return GraphTextTrainer(config, **kwargs)

    def _init_optimizer(self):
        """Initialize optimizer on the model's (FSDP-sharded) parameters.