Diff Coverage

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

Source File Diff Coverage (%) Missing Lines
hyper_parallel/core/dtensor/layout.py 100%  
hyper_parallel/core/fully_shard/utils.py 100%  
hyper_parallel/core/optimizer/__init__.py 100%  
hyper_parallel/core/shard/utils.py 50.0% 123
hyper_parallel/distributed/_builder/fsdp_adapter.py 0.0% 312-313,318
hyper_parallel/distributed/expert_parallel/experts.py 40.9% 103-109,131-132,159,162-164
hyper_parallel/distributed/mesh.py 5.9% 188,198-202,206-208,212-213,218,223,229-231
hyper_parallel/trainer/base.py 33.3% 496-497
hyper_parallel/trainer/callbacks/environ_meter_callback.py 87.5% 93,102,176,183,187
hyper_parallel/trainer/callbacks/tqdm_callback.py 100%  
hyper_parallel/trainer/text_trainer.py 100%  
hyper_parallel/core/shard/utils.py
119
120
121
122
123
124
125
126
127
        label_smoothing: float = 0.0,
) -> Tensor:
    """Distributed cross_entropy entry used by shard dispatch."""
    # Defer the components import to preserve the lightweight models import boundary.
    from hyper_parallel.components.losses._vocab_parallel_cross_entropy import (  # pylint: disable=C0415
        DistributedCrossEntropyFunction,
    )

    input_dtensor = None
hyper_parallel/distributed/_builder/fsdp_adapter.py
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
            aliases_by_parameter[parameter].append(parameter_fqn)

        owner_by_parameter = {}
        for parameter, parameter_fqns in aliases_by_parameter.items():
            common_wrap_modules = []
            for wrap_module in wrap_modules:
                if all(
                    parameter_fqn.startswith(f"{wrap_module.fqn}.")
                    for parameter_fqn in parameter_fqns
                ):
                    common_wrap_modules.append(wrap_module)
            if common_wrap_modules:
                owner_by_parameter[parameter] = max(
                    common_wrap_modules,
                    key=lambda wrap_module: wrap_module.fqn.count("."),
hyper_parallel/distributed/expert_parallel/experts.py
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113


def _run_grouped_swiglu_expert_forward(experts, sorted_states, local_expert_counts):
    """Run the grouped GEMM implementation for sorted expert tokens."""
    grouped_forward = getattr(experts, "forward_expert_major", None)
    if callable(grouped_forward):
        return grouped_forward(sorted_states, local_expert_counts)
    gate_weight, up_weight, down_weight = resolve_swiglu_weights(experts)
    if up_weight is not None:
        raise ValueError("EP grouped GEMM currently requires packed gate_up_proj weights")
    return npu_grouped_swiglu(
        sorted_states,
        gate_weight,
        down_weight,
        local_expert_counts,
127
128
129
130
131
132
133
134
135
136
                expert_states,
                gate_weight[local_expert_index],
            ).chunk(2, dim=-1)
        else:
            gate_states = F.linear(expert_states, gate_weight[local_expert_index])
            up_states = F.linear(expert_states, up_weight[local_expert_index])
        activation = getattr(experts, "ep_act_fn", F.silu)
        sorted_outputs.append(
            F.linear(  # pylint: disable=not-callable
                activation(gate_states) * up_states,
155
156
157
158
159
160
161
162
163
164
165
166
167
168
        local_expert_indices,
        minlength=experts.local_expert_count,
    )
    if getattr(experts, "ep_use_grouped_gemm", False):
        sorted_output = _run_grouped_swiglu_expert_forward(
            experts, sorted_states, local_expert_counts
        )
        output = torch.empty_like(sorted_output)
        output[token_order] = sorted_output
        return output
    sorted_output = _run_per_expert_swiglu_forward(
        experts, sorted_states, local_expert_counts
    )
    output = torch.empty_like(sorted_output)
hyper_parallel/distributed/mesh.py
184
185
186
187
188
189
190
191
192
            dense_mesh_kwargs["rank_list"] = dense_rank_list
        self.fsdp_non_moe_mesh = _init_topology_mesh(dense_mesh_kwargs)
        _validate_dense_tp_rank_layout(self.device_mesh, self.fsdp_non_moe_mesh)

        self._build_expert_parallel_mesh(device_type, init_backend, dense_rank_list)

    def _build_expert_parallel_mesh(
        self,
        device_type: str,
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
        dense_rank_list: tuple[int, ...] | None,
    ) -> None:
        """Build the expert-parallel FSDP mesh when expert parallelism is enabled."""
        self.fsdp_moe_mesh = None
        if self.ep_size <= 1:
            return
        expert_domain_size = math.prod(self.device_mesh.mesh_shape)
        if expert_domain_size % self.ep_size != 0:
            raise ValueError(
                f"expert domain size ({expert_domain_size}) must be divisible by "
                f"ep_size ({self.ep_size})"
            )
        edp_size = expert_domain_size // self.ep_size
        if edp_size % self.edp_shard_size != 0:
            raise ValueError(
                f"expert data-parallel size ({edp_size}) must be divisible by "
                f"edp_shard_size ({self.edp_shard_size})"
            )
        edp_replicate_size = edp_size // self.edp_shard_size
        expert_mesh_shape = (
            (edp_replicate_size, self.edp_shard_size, self.ep_size)
            if edp_replicate_size > 1
            else (self.edp_shard_size, self.ep_size)
        )
        expert_mesh_names = (
            ("edp_replicate", "edp_shard", "ep")
            if edp_replicate_size > 1
            else ("edp_shard", "ep")
        )
        expert_mesh_kwargs = {
            "device_type": device_type,
            "mesh_shape": expert_mesh_shape,
            "mesh_dim_names": expert_mesh_names,
            "init_backend": init_backend,
225
226
227
228
229
230
231
232
233
234
235
            "mesh_shape": expert_mesh_shape,
            "mesh_dim_names": expert_mesh_names,
            "init_backend": init_backend,
        }
        if init_backend and dense_rank_list is not None:
            expert_mesh_kwargs["rank_list"] = dense_rank_list
        self.fsdp_moe_mesh = _init_topology_mesh(expert_mesh_kwargs)


def _validate_dense_tp_rank_layout(device_mesh: Any, fsdp_non_moe_mesh: Any) -> None:
    """Ensure both TP child meshes describe the same groups and local rank."""
hyper_parallel/trainer/base.py
492
493
494
495
496
497
498
499
500
501
        Args:
            micro_batch: Prepared inputs for the current micro step.
            **kwargs: Additional callback context.
        """
        for callback in self._callbacks:
            callback.on_micro_step_begin(self.state, micro_batch, **kwargs)

    def on_step_end(
        self,
        loss: Optional[float] = None,
hyper_parallel/trainer/callbacks/environ_meter_callback.py
89
90
91
92
93
94
95
96
97
            return token_count

        labels = batch.get("labels")
        if labels is not None and callable(getattr(labels, "sum", None)):
            return (labels != IGNORE_INDEX).sum()

        attention_mask = batch.get("attention_mask")
        attention_mask_shape = getattr(attention_mask, "shape", ())
        if (
 98
 99
100
101
102
103
104
105
106
            len(attention_mask_shape) <= 2
            and attention_mask is not None
            and callable(getattr(attention_mask, "sum", None))
        ):
            return attention_mask.sum()

        input_ids = batch.get("input_ids")
        input_numel = cls._tensor_numel(input_ids)
        if input_numel is not None:
172
173
174
175
176
177
178
179
180
                if callable(getattr(token_count, "clone", None)):
                    token_count = token_count.clone()
                self._local_step_tokens = token_count
            else:
                self._local_step_tokens = self._local_step_tokens + token_count
            self._local_step_samples += self._batch_samples(batch)

    def _global_samples(self) -> int:
        """Reduce samples across DP+CP while removing CP replicas."""
179
180
181
182
183
184
185
186
187
188
189
190
191
    def _global_samples(self) -> int:
        """Reduce samples across DP+CP while removing CP replicas."""
        cp_size = int(getattr(self.trainer.mesh, "cp_size", 1))
        if cp_size < 1:
            raise ValueError(f"mesh.cp_size must be positive, but got {cp_size}")
        reduced_samples = self._reduce(self._local_step_samples, op="sum")
        global_samples = reduced_samples / cp_size
        if not global_samples.is_integer():
            raise ValueError(
                "Reduced sample count must be divisible by cp_size, "
                f"but got reduced_samples={reduced_samples} and cp_size={cp_size}"
            )
        return int(global_samples)