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/compiler.py 100%  
hyper_parallel/compile/examples/automodel_tp/train.py 0.0% 277-278
hyper_parallel/compile/graph_parallel_plan.py 94.1% 187,193,234,283-284
hyper_parallel/compile/passes/parallel/fsdp_pass.py 100%  
hyper_parallel/compile/passes/parallel/pp_pass.py 100%  
hyper_parallel/compile/passes/pipeline.py 83.3% 85
hyper_parallel/compile/trainer.py 100%  
hyper_parallel/components/functional/npu_fusion_attention.py 73.0% 96,160,213-214,245-246,249,303,336-337
hyper_parallel/components/functional/rotary_embedding.py 0.0% 95,98
hyper_parallel/components/modules/mla_attention.py 73.3% 215-217,250
hyper_parallel/core/shard/ops/parallel_argsort.py 100%  
hyper_parallel/core/shard/ops/parallel_cumsum.py 100%  
hyper_parallel/core/shard/ops/parallel_embedding.py 100%  
hyper_parallel/core/shard/ops/parallel_rotary_position_embedding.py 100%  
hyper_parallel/core/shard/ops/parallel_sort.py 100%  
hyper_parallel/core/shard/utils.py 50.0% 123
hyper_parallel/core/utils/communication.py 35.3% 207-211,224,227-230,233
hyper_parallel/distributed/context_parallel/mla_context_parallel.py 75.0% 54,56,85,87,89,91,115,129-131,143,147,151-152,162-163,170,183,190,200,202,211,213,229,277,280,290,293,298,300,303-306,315,346,349,356,362,365,373,376-378,382-389,391-392
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/compile/examples/automodel_tp/train.py
273
274
275
276
277
278
279
280
281
282
            "attention, no patch needed"
        )

    # 4. FSDP plan for the graph-mode FSDPPass (mark everything).
    fsdp_plan = GraphParallelPlan()
    fsdp_plan.fsdp_mark_pattern("*")

    # 5. GraphTrainer reuses the automodel mesh (tp group already created);
    #    it registers the dp sub-mesh as "fsdp" and back-fills fsdp_degree.
    tcfg = cfg["train"]
hyper_parallel/compile/graph_parallel_plan.py
183
184
185
186
187
188
189
190
191
        raise ValueError("Must provide either config_path or model_name")

    if config_path is None:
        if not model_name or not isinstance(model_name, str):
            raise ValueError("model_name must be a non-empty string")
        if ".." in model_name or "/" in model_name or "\\" in model_name:
            raise ValueError(
                f"Invalid model_name '{model_name}': must not contain path "
                "separators or parent directory references"
189
190
191
192
193
194
195
196
197
            raise ValueError(
                f"Invalid model_name '{model_name}': must not contain path "
                "separators or parent directory references"
            )
        config_path = DEFAULT_CONFIG_DIR / model_name / "config.yaml"
    else:
        config_path = Path(config_path)

    if not config_path.exists():
230
231
232
233
234
235
236
237
238
    section = config.get(key) or {}
    if not isinstance(section, dict):
        # A present-but-empty section parses to None and is normalized to {}
        # above, so anything landing here is a real scalar/sequence typo.
        raise ValueError(
            f"YAML '{key}' section must be a mapping (e.g. nested keys or an "
            f"empty section); got {type(section).__name__} in {config_path}"
        )
    return section
279
280
281
282
283
284
285
286
287
288
        if isinstance(stage, dict):
            stage_idx = stage.get("stage", idx)
            module_fqns = list(stage.get("modules", []))
        else:
            stage_idx = idx
            module_fqns = list(stage)
        plan.pp_stage(stage_idx, module_fqns)


def create_all_fsdp_plan() -> GraphParallelPlan:
hyper_parallel/compile/passes/pipeline.py
81
82
83
84
85
86
87
88
89
        # 2. Execution layer: Parallel dimension partitioning
        if getattr(self.config, "fsdp_enabled", False):
            self.passes.append(FSDPPass(parallel_plan=self.parallel_plan))
        if getattr(self.config, "pp_enabled", False):
            self.passes.append(PpPass(parallel_plan=self.parallel_plan))

        # 3. Communication-compute overlap optimization
        if getattr(self.config, "enable_overlap", False):
            self.passes.append(AutoOverlapPass())
hyper_parallel/components/functional/npu_fusion_attention.py
 92
 93
 94
 95
 96
 97
 98
 99
100
) -> list[int]:
    """Normalize one cumulative-length representation."""
    if isinstance(value, torch.Tensor):
        if value.ndim != 1:
            raise ValueError("packed cumulative sequence lengths must be one-dimensional")
        value = value.tolist()
    if any(isinstance(item, bool) or not isinstance(item, Integral) for item in value):
        raise ValueError("packed cumulative sequence lengths must be integers")
    lengths = [int(item) for item in value]
156
157
158
159
160
161
162
163
164
    key_lengths = _coalesce_lengths(
        kwargs, _KEY_LENGTH_ALIASES, name="key/value sequence lengths",
    )
    if kwargs.get("packed_seq_params") is not None and query_lengths is None and key_lengths is None:
        raise ValueError("packed_seq_params must provide cumulative query and key/value sequence lengths")
    if (query_lengths is None) != (key_lengths is None):
        raise ValueError("packed attention requires both query and key/value lengths")
    if query_lengths is not None:
        if query_lengths[-1] != query_tokens:
209
210
211
212
213
214
215
216
217
218
            mask = torch.ones((2048, 2048), dtype=torch.bool, device=query.device).triu(diagonal=1)
            return query, key, value, "BNSD", mask, 3
        return query, key, value, "BNSD", _causal_attention_mask(query, key, sliding_window), 0
    npu_mask = None if attention_mask is None else _npu_attention_mask(attention_mask)
    if npu_mask is not None and is_causal and sparse_mode == 0:
        npu_mask = npu_mask | _causal_attention_mask(query, key, sliding_window)
    return query, key, value, "BNSD", npu_mask, sparse_mode


def _prepare_fusion_attention_context(
241
242
243
244
245
246
247
248
249
250
251
252
253
        sparse_mode=options[2],
    )
    valid_rows = None
    if query_lengths is None and prepared[4] is not None and prepared[5] == 0:
        mask = prepared[4]
        valid_rows = (~mask).any(dim=-1)
        # Give fully masked rows one finite softmax entry, then suppress their
        # output and upstream gradient. Some NPU kernels otherwise return nonzero rows.
        mask[..., 0] &= valid_rows
    return _FusionAttentionContext(
        *prepared[:5],
        prepared[5],
        query_lengths,
299
300
301
302
303
304
305
306
307
        ``True`` for positions that participate in attention; an additive mask
        uses zero for those positions.
    """
    # Length parsing is also used by CPU CP validation; load the optional kernel only here.
    import torch_npu  # pylint: disable=C0415

    if kwargs.get("indices") is not None:
        raise ValueError(
            "npu_fusion_attention_forward does not consume sparse attention indices; "
332
333
334
335
336
337
338
339
340
341
        sparse_mode=context.sparse_mode,
        actual_seq_qlen=context.query_lengths,
        actual_seq_kvlen=context.key_lengths,
    )[0]
    if context.valid_rows is not None:
        output = output.masked_fill(~context.valid_rows.unsqueeze(-1), 0.0)
    if context.is_packed:
        output = output.reshape(
            context.batch_size,
            context.query_length,
hyper_parallel/components/functional/rotary_embedding.py
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
            rotary_mode="interleave",
        )
    rotated = rotated.permute(1, 2, 0, 3)
    # Transpose pairs instead of strided slices to avoid scattered writes in backward.
    rotated = rotated.unflatten(-1, (-1, 2)).transpose(-1, -2).flatten(-2)
    if unsqueeze_dim == 2:
        rotated = rotated.transpose(1, 2)
    return torch.cat((rotated, pass_through), dim=-1) if pass_through.shape[-1] else rotated


def apply_rotary_pos_emb_interleave(
    q: torch.Tensor,
hyper_parallel/components/modules/mla_attention.py
211
212
213
214
215
216
217
218
219
220
221
        return self.q_a_layernorm(query_latent), self.kv_a_layernorm(kv_nope), key_rope

    def _project_latents(self, hidden_states: torch.Tensor) -> _MLALatents:
        """Project hidden states into query and compressed KV latent states."""
        batch_size, sequence_length = hidden_states.shape[:-1]
        query_latent, kv_nope, key_rope = self._normalized_latents(hidden_states)
        query_states = self.q_b_proj(query_latent).view(
            batch_size,
            sequence_length,
            self.num_heads,
            self.qk_head_dim,
246
247
248
249
250
251
252
253
254
        latents = self._project_latents(hidden_states)
        query_rope, key_rope = latents.query_rope, latents.key_rope
        if position_embeddings is not None:
            # The original NPU RoPE implementation is optional on CPU CP users.
            from hyper_parallel.components.functional import (  # pylint: disable=C0415
                apply_rotary_pos_emb, apply_rotary_pos_emb_interleave,
            )

            cos, sin = position_embeddings
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/core/utils/communication.py
203
204
205
206
207
208
209
210
211
212
213
214
215

        Returns:
            One tensor per process-group rank, in group-rank order.
        """
        ctx.group = group
        tensor = tensor.contiguous()
        outputs = [torch.empty_like(tensor) for _ in range(dist.get_world_size(group))]
        dist.all_gather(outputs, tensor, group=group)
        return tuple(outputs)

    @staticmethod
    def backward(ctx: Any, *gradients: Tensor) -> tuple[Tensor, None]:
        """Sum all consumers' contributions to the local input shard.
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237

        Returns:
            The summed gradient for this rank's input and no group gradient.
        """
        if dist.get_backend(ctx.group) == "gloo":
            # Keep older Gloo versions supported and avoid upstream AllToAll
            # emulation that treats subgroup source ranks as global scatter ranks.
            gradients = dist_func.all_reduce(torch.stack(gradients), op=dist.ReduceOp.SUM, group=ctx.group)
            return gradients[dist.get_rank(ctx.group)], None
        output = torch.empty_like(gradients[0])
        gradient = dist_func.reduce_scatter(
            output, [part.contiguous() for part in gradients], op=dist.ReduceOp.SUM, group=ctx.group,
        )
        return gradient, None


class _AsyncA2ALazyBwd(torch.autograd.Function):
    """All-to-all whose forward AND backward return ``AsyncCollectiveTensor``.
hyper_parallel/distributed/context_parallel/mla_context_parallel.py
50
51
52
53
54
55
56
57
58
59
60

    def __post_init__(self) -> None:
        for name, value in vars(self).items():
            if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
                raise ValueError(f"{name} must be a positive integer, got {value}")
        if self.rope_dim % 2:
            raise ValueError("rope_dim must be even")

    @property
    def qk_dim(self) -> int:
        """Return the score dimension, including the decoupled RoPE band."""
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
    tp_degree: int = 1

    def __post_init__(self) -> None:
        if not isinstance(self.cp_degree, int) or isinstance(self.cp_degree, bool) or self.cp_degree < 1:
            raise ValueError("cp_degree must be a positive integer")
        if self.strategy not in MLA_CP_STRATEGIES:
            raise ValueError(f"Unknown MLA CP strategy: {self.strategy}")
        if self.backend not in ("npu", "sdpa"):
            raise ValueError(f"Unknown MLA CP backend: {self.backend}")
        if not isinstance(self.tp_degree, int) or isinstance(self.tp_degree, bool) or self.tp_degree < 1:
            raise ValueError("tp_degree must be a positive integer")
        if self.dimensions.heads % (self.tp_degree * self.cp_degree):
            raise ValueError("MLA heads must be divisible by TP degree times CP degree")

    @property
111
112
113
114
115
116
117
118
119
        Raises:
            ValueError: If the scale is nonfinite or nonpositive.
        """
        if not math.isfinite(scale) or scale <= 0:
            raise ValueError("MLA attention scaling must be finite and positive")


def mla_all_gather(tensor: torch.Tensor, sequence_dim: int, cp_mesh: Any) -> torch.Tensor:
    """Gather source tokens and sum all consumer gradients back to their owner.
125
126
127
128
129
130
131
132
133
134
135

    Returns:
        Full-sequence tensor in CP group-rank order.
    """
    if cp_mesh is None or cp_mesh.size() == 1:
        return tensor
    return differentiable_all_gather_concat(tensor, cp_mesh.get_group(), cp_mesh.size(), sequence_dim)


def _apply_mla_rope(
    query: torch.Tensor,
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
    interleaved: bool = False,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Rotate local BSND Q and shared K with the existing NPU kernels or CPU math."""
    if position_embeddings is None:
        return query, key
    cos, sin = (part.to(query.dtype) for part in position_embeddings)
    if query.device.type == "npu":
        # The shared NPU RoPE module is optional for CPU reference execution.
        from hyper_parallel.components.functional.rotary_embedding import (  # pylint: disable=C0415
            apply_rotary_pos_emb, apply_rotary_pos_emb_interleave,
        )

        rope = apply_rotary_pos_emb_interleave if interleaved else apply_rotary_pos_emb
        return rope(query, key, cos, sin, unsqueeze_dim=2)
    cos, sin = cos.unsqueeze(2), sin.unsqueeze(2)
    outputs = []
    for tensor in (query, key):
        if interleaved:
158
159
160
161
162
163
164
165
166
167
            half_cos, half_sin = cos[..., :tensor.shape[-1] // 2], sin[..., :tensor.shape[-1] // 2]
            outputs.append(torch.cat((real * half_cos - imaginary * half_sin,
                                      imaginary * half_cos + real * half_sin), dim=-1))
        else:
            first, second = tensor.chunk(2, dim=-1)
            outputs.append(tensor * cos + torch.cat((-second, first), dim=-1) * sin)
    return tuple(outputs)


def _run_attention(module, query, key, value, *, backend, lengths, causal, attention_mask):
166
167
168
169
170
171
172
173
174

def _run_attention(module, query, key, value, *, backend, lengths, causal, attention_mask):
    """Use the shared FA interface after CP has restored full Q/K token order."""
    if backend == "npu":
        return npu_fusion_attention_forward(
            module, query, key, value, attention_mask, scaling=module.scaling,
            is_causal=causal, actual_seq_len=lengths,
            pre_tokens=2147483647, next_tokens=0 if causal else 2147483647,
        )[0]
179
180
181
182
183
184
185
186
187
        return F.scaled_dot_product_attention(  # pylint: disable=not-callable
            query, key, value, attn_mask=allowed, dropout_p=0.0,
            is_causal=causal and allowed is None, scale=module.scaling,
        ).transpose(1, 2)
    outputs = [
        F.scaled_dot_product_attention(  # pylint: disable=not-callable
            query[:, :, begin:end], key[:, :, begin:end], value[:, :, begin:end],
            dropout_p=0.0, is_causal=causal, scale=module.scaling,
        ).transpose(1, 2)
186
187
188
189
190
191
192
193
194
            dropout_p=0.0, is_causal=causal, scale=module.scaling,
        ).transpose(1, 2)
        for begin, end in zip([0] + lengths[:-1], lengths)
    ]
    return torch.cat(outputs, dim=1)


def _positions(
    embeddings: tuple[torch.Tensor, torch.Tensor] | None,
196
197
198
199
200
201
202
203
204
205
206
    rope_dim: int,
) -> tuple[torch.Tensor, torch.Tensor] | None:
    """Validate local-token RoPE frequencies and broadcast only the batch axis."""
    if embeddings is None:
        return None
    if not isinstance(embeddings, (tuple, list)) or len(embeddings) != 2:
        raise ValueError("position_embeddings must be a local (cos, sin) pair")
    batch, sequence = state.query_latent.shape[:2]
    parts = []
    for tensor in embeddings:
        if tensor.ndim == 2:
207
208
209
210
211
212
213
214
215
216
            tensor = tensor.unsqueeze(0)
        if tensor.shape not in ((1, sequence, rope_dim), (batch, sequence, rope_dim)):
            raise ValueError("RoPE cos/sin must describe local tokens with shape [B, S_local, rope_dim]")
        if tensor.requires_grad:
            raise ValueError("MLA CP currently requires non-trainable RoPE frequencies")
        if tensor.device != state.query_latent.device:
            raise ValueError("RoPE and MLA activations must be on the same device")
        parts.append(tensor.expand(batch, -1, -1))
    return tuple(parts)

225
226
227
228
229
230
231
232
233
    """Require global boolean masks and reject unsupported packed-mask combinations."""
    if mask is None:
        return
    if lengths is not None:
        raise ValueError("packed MLA CP currently accepts document causal/noncausal masks only")
    if mask.dtype != torch.bool:
        raise ValueError("MLA CP masks must be boolean (True=allowed); additive bias is not supported")
    valid = (
        mask.ndim == 2 and mask.shape == (sequence, sequence)
273
274
275
276
277
278
279
280
281
282
283
def _validate_forward_kwargs(kwargs, position_embeddings):
    """Validate forwarded model options before routing any latent."""
    position_ids = kwargs.pop("position_ids", None)
    if position_ids is not None and position_embeddings is None:
        raise ValueError("position_ids require precomputed local position_embeddings in MLA CP")
    kwargs.pop("cache_position", None)
    if kwargs.pop("use_cache", False) or kwargs.pop("output_attentions", False):
        raise ValueError("MLA CP does not return KV caches or attention weights")
    if kwargs:
        raise ValueError(f"Unsupported MLA CP forward arguments: {sorted(kwargs)}")

286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
    """Check latent shape, dtype and backend requirements before collectives."""
    dimensions = plan.dimensions
    query_latent, kv_latent, key_rope = state.query_latent, state.kv_latent, state.key_rope
    if query_latent.ndim != 3:
        raise ValueError("MLA CP latents must have BSR layout")
    batch, sequence = query_latent.shape[:2]
    if batch < 1 or sequence < 1:
        raise ValueError("MLA CP requires nonempty batches and sequence shards")
    for tensor, width in ((query_latent, dimensions.q_rank), (kv_latent, dimensions.kv_rank),
                          (key_rope, dimensions.rope_dim)):
        if (tensor.shape != (batch, sequence, width) or tensor.device != query_latent.device
                or tensor.dtype != query_latent.dtype):
            raise ValueError("MLA CP latent shapes/devices/dtypes do not match the plan")
    if module.training and module.attention_dropout != 0:
        raise ValueError("MLA CP requires attention dropout=0")
    plan.validate_scale(module.scaling)
    if plan.backend == "npu":
        if query_latent.device.type != "npu" or query_latent.dtype not in (torch.float16, torch.bfloat16):
            raise ValueError("NPU MLA CP requires NPU BF16/FP16 activations")
        if dimensions.value_dim > dimensions.qk_dim:
            raise ValueError("NPU MLA CP requires value_dim <= qk_dim")


class MLAContextParallel:
    """Run CP over normalized latents; parameter-gradient reduction belongs to the trainer."""
311
312
313
314
315
316
317
318
319

    def __init__(self, plan: MLACPPlan, cp_mesh: Any) -> None:
        """Bind a validated execution plan to the framework-owned CP mesh."""
        if (1 if cp_mesh is None else cp_mesh.size()) != plan.cp_degree:
            raise ValueError("MLACPPlan degree must match the CP mesh")
        self.plan = plan
        self.cp_mesh = cp_mesh

    def __call__(
342
343
344
345
346
347
348
349
350
351
352
353
        )
        if lengths is not None and batch != 1:
            raise ValueError("packed MLA CP currently requires batch_size=1")
        if lengths != key_lengths:
            raise ValueError("MLA Ulysses requires matching global query and key document boundaries")
        causal = getattr(module, "is_causal", True) if is_causal is None else is_causal
        if not isinstance(causal, bool):
            raise ValueError("is_causal must be bool")
        _validate_mask(attention_mask, lengths, batch, sequence * degree, query_latent.device)
        positions = _positions(position_embeddings, state, dimensions.rope_dim)
        query, rotated_key = _query_and_key_rope(module, state, positions, self.plan.local_heads)
        if self.plan.strategy == "expanded_ulysses" or degree == 1:
352
353
354
355
356
357
358
359
360
        query, rotated_key = _query_and_key_rope(module, state, positions, self.plan.local_heads)
        if self.plan.strategy == "expanded_ulysses" or degree == 1:
            query, key, value = self._expanded(module, state, query, rotated_key)
        else:
            query, key, value = self._latent(module, state, query, rotated_key)
        output = _run_attention(
            module, query, key, value, backend=self.plan.backend, lengths=lengths,
            causal=causal, attention_mask=attention_mask,
        )
358
359
360
361
362
363
364
365
366
367
368
369
            module, query, key, value, backend=self.plan.backend, lengths=lengths,
            causal=causal, attention_mask=attention_mask,
        )
        if degree > 1:
            output = ulysses_head_to_seq(output, 1, 2, self.cp_mesh)
        expected = (batch, sequence, self.plan.local_heads, dimensions.value_dim)
        if tuple(output.shape) != expected:
            raise RuntimeError(f"MLA CP output has shape {tuple(output.shape)}, expected {expected}")
        return output, None

    def _expanded(self, module, state, query, key_rope):
        """Pack unequal-width QKV into one sequence-to-head exchange."""
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
        """Pack unequal-width QKV into one sequence-to-head exchange."""
        key, value = _key_value(module, state.kv_latent, key_rope, self.plan.local_heads, None)
        if self.plan.cp_degree == 1:
            return query, key, value
        widths = (query.shape[-1], key.shape[-1], value.shape[-1])
        # Keep projection outputs in BSND while packing to reduce layout copies.
        # Source-rank/sequence reconstruction is also a view when B=1.
        payload = torch.cat(tuple(tensor.transpose(1, 2) for tensor in (query, key, value)), dim=-1)
        payload = ulysses_seq_to_head(payload, 1, 2, self.cp_mesh)
        return tuple(tensor.transpose(1, 2) for tensor in payload.split(widths, dim=-1))

    def _latent(self, module, state, query, key_rope):
        """Exchange Q heads, gather shared KV latents, then expand only owned heads."""
        dimensions = self.plan.dimensions
        heads = self.plan.compute_heads
        rank = 0 if self.cp_mesh is None else self.cp_mesh.get_local_rank()
        start = rank * heads
        if self.plan.cp_degree > 1:
            query = ulysses_seq_to_head(query, 2, 1, self.cp_mesh)
        payload = mla_all_gather(torch.cat((state.kv_latent, key_rope), dim=-1), 1, self.cp_mesh)
        kv_latent, key_rope = payload.split((dimensions.kv_rank, dimensions.rope_dim), dim=-1)
        # A contiguous latent preserves the native Linear bias-fusion path after payload splitting.
        key, value = _key_value(module, kv_latent.contiguous(), key_rope, heads, (start, start + heads))
        return query, key, value
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)