Diff Coverage

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

Source File Diff Coverage (%) Missing Lines
hyper_parallel/core/activation_memory/_backend.py 100%  
hyper_parallel/core/context_parallel/context_parallel.py 100%  
hyper_parallel/core/context_parallel/dsa_context_parallel.py 100%  
hyper_parallel/core/context_parallel/utils.py 100%  
hyper_parallel/core/dtensor/_utils.py 100%  
hyper_parallel/core/expert_parallel/expert_parallel.py 100%  
hyper_parallel/core/fully_shard/api.py 33.3% 169,294,340,356,372,383
hyper_parallel/core/fully_shard/hsdp_scheduler.py 0.0% 260
hyper_parallel/core/fully_shard/utils.py 100%  
hyper_parallel/core/shard/api.py 100%  
hyper_parallel/core/shard/ops/parallel_gather.py 100%  
hyper_parallel/core/shard/ops/parallel_npu_flash_attention_score.py 33.3% 1295,1313
hyper_parallel/core/shard/ops/parallel_scaled_dot_product_attention.py 66.7% 186
hyper_parallel/core/shard/utils.py 33.3% 46-55
hyper_parallel/core/utils/__init__.py 100%  
hyper_parallel/core/utils/communication.py 48.2% 43-49,64-68,78,95,107-110,123-125,127-134,136,141-142,167,237-241,245-247,252-254,258-260,269-275,277-281,287-291,293-295,297-302,307-310,315-318,320-323,331-335,361-364,408-410,432,439,456-458
hyper_parallel/core/utils/moe_utils.py 100%  
hyper_parallel/distributed/_builder/tp_collective_lowering.py 20.0% 44,51,57,65
hyper_parallel/distributed/context_parallel/attention.py 50.0% 305
hyper_parallel/distributed/context_parallel/collectives.py 50.0% 470
hyper_parallel/distributed/context_parallel/gated_delta_net.py 33.3% 164,910
hyper_parallel/distributed/context_parallel/kimi_delta_attention.py 33.3% 193,316
hyper_parallel/core/fully_shard/api.py
165
166
167
168
169
170
171
172
173
            raise ValueError(f"requires_grad_sync must be bool but got {requires_grad_sync}.")
        if not hasattr(self, "hsdp_scheduler"):
            raise ValueError("call hsdp interface first.")

        for _, module in self.named_modules():
            if isinstance(module, HSDPModule):
                module.hsdp_scheduler.set_requires_grad_sync(requires_grad_sync)

    def zero_grad(self):
290
291
292
293
294
295
296
297
        # Pass an IncompatibleKeys with the same attribute names as PyTorch
        # so external hooks can safely read .missing_keys/.unexpected_keys.
        _IK = namedtuple("IncompatibleKeys", ["missing_keys", "unexpected_keys"])
        incompatible_keys = _IK([], [])
        for _, module in self_module.named_modules():
            hooks = module._load_state_dict_post_hooks  # pylint: disable=protected-access
            for hook in hooks.values():
                hook(module, incompatible_keys)
336
337
338
339
340
341
342
343
344
                "Currently impl is equal to recurse=True, "
                "need support module_param mapping."
            )
        self_module = cast(nn.Module, self)
        for _, module in self_module.named_modules():
            if isinstance(module, HSDPModule):
                module.hsdp_scheduler.set_requires_all_reduce(requires_all_reduce)

    def set_reshard_after_forward(self, reshard_after_forward: bool, recurse: bool = True) -> None:
352
353
354
355
356
357
358
359
360
                "Currently impl is equal to recurse=True, "
                "need support module_param mapping."
            )
        self_module = cast(nn.Module, self)
        for _, module in self_module.named_modules():
            if isinstance(module, HSDPModule):
                module.hsdp_scheduler.set_reshard_after_forward(reshard_after_forward)

    def set_reshard_after_backward(self, reshard_after_backward: bool, recurse: bool = True) -> None:
368
369
370
371
372
373
374
375
376
                "Currently impl is equal to recurse=True, "
                "need support module_param mapping."
            )
        self_module = cast(nn.Module, self)
        for _, module in self_module.named_modules():
            if isinstance(module, HSDPModule):
                module.hsdp_scheduler.set_reshard_after_backward(reshard_after_backward)

    def set_reduce_op_type(self, reduce_op_type, recurse: bool = True) -> None:
379
380
381
382
383
384
385
386
387
        support reduce_op_type "avg" and "sum", default is "avg"
        """
        self_module = cast(nn.Module, self)
        if recurse:
            sub_modules = [m for _, m in self_module.named_modules()]
        else:
            sub_modules = [self_module]
        for module in sub_modules:
            if isinstance(module, HSDPModule):
hyper_parallel/core/fully_shard/hsdp_scheduler.py
256
257
258
259
260
261
262
263
264
        if self.scheduler_ctx.root_module is None:
            tree_ctx = self.scheduler_ctx
            tree_ctx.root_module = self.cell
            registered_schedulers = set()
            for module_name, module in tree_ctx.root_module.named_modules():
                from hyper_parallel.core.fully_shard.api import HSDPModule  # pylint: disable=C0415
                if isinstance(module, HSDPModule):
                    submod_scheduler = module.hsdp_scheduler
                    if submod_scheduler is None or id(submod_scheduler) in registered_schedulers:
hyper_parallel/core/shard/ops/parallel_npu_flash_attention_score.py
1291
1292
1293
1294
1295
1296
1297
1298
1299
        if dim_map == "None":
            return 0

        if isinstance(dim_map, str):
            rank = dist.get_rank()
            rank_list = layout.mesh.get_rank_list_along_axis(dim_map)
            if rank in rank_list:
                return rank_list.index(rank)
            return 0
1309
1310
1311
1312
1313
1314
1315
1316
                    f"Seq dim is sharded by multiple axes {non_none_axes}. "
                    f"Using the last axis for split_id calculation."
                )
            axis_name = non_none_axes[-1]
            rank = dist.get_rank()
            rank_list = layout.mesh.get_rank_list_along_axis(axis_name)
            if rank in rank_list:
                return rank_list.index(rank)
hyper_parallel/core/shard/ops/parallel_scaled_dot_product_attention.py
182
183
184
185
186
187
188
189
                    f"Seq dim is sharded by multiple axes {non_none_axes}. "
                    f"Using the last axis for split_id calculation."
                )
            axis_name = non_none_axes[-1]
            rank = dist.get_rank()
            rank_list = layout.mesh.get_rank_list_along_axis(axis_name)
            if rank in rank_list:
                return rank_list.index(rank)
hyper_parallel/core/shard/utils.py
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
def get_op_name(func):
    """Return the registry name for a Torch callable or operator overload."""
    if hasattr(func, "__name__"):
        return func.__name__
    if isinstance(func, OpOverload):
        return func.name.split("::")[-1].split(".")[0]
    if isinstance(func, OpOverloadPacket):
        return func.name.split("::")[-1]
    func_str = str(func)
    if "built-in function" in func_str:
        return func_str.split()[-1].strip(">")
    if "function" in func_str:
        return func_str.split()[1]
    return "unknown_op"


def get_cell_construct(cell):
    """Return the Torch module forward callable."""
hyper_parallel/core/utils/communication.py
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53

    Returns:
        int: The rank of the current process inside ``group``.
    """
    if not torch.distributed.is_available() or not torch.distributed.is_initialized():
        return 0
    if group is None:
        return torch.distributed.get_rank()
    if hasattr(group, "rank"):
        return group.rank()
    return torch.distributed.get_group_rank(group, torch.distributed.get_rank())


def get_device_handle(device_type: str = "npu"):
    """Return the ``torch`` device module for ``device_type``.
60
61
62
63
64
65
66
67
68
69
70
71
72

    Raises:
        RuntimeError: If torch exposes no module for ``device_type``.
    """
    try:
        handle = getattr(torch, device_type)
    except AttributeError as e:
        raise RuntimeError(f"Failed to resolve device handle: 'torch.{device_type}'.") from e
    return handle


def get_world_size() -> int:
    """Return the Torch distributed world size.
74
75
76
77
78
79
80
81
82
    Kept as a thin wrapper so single-process call sites and white-box unit
    tests share one patchable seam; ``torch.distributed.get_world_size`` itself
    raises before ``init_process_group``.
    """
    return dist.get_world_size()


# Process-group cache keyed by the string form of the sorted rank tuple, e.g.
# ``"(0, 1, 2, 3)"``. Lives here rather than in a module-specific ``utils`` so
91
92
93
94
95
96
97
98
    Under ``torch.compile`` the contiguity of a graph input is not knowable at
    trace time, so the copy is emitted unconditionally.
    """
    if torch.compiler.is_compiling():
        return x.contiguous()
    if not x.is_contiguous() or x.storage_offset() != 0:
        return x.contiguous()
    return x
103
104
105
106
107
108
109
110
111
112
113
114
# ---------------------------------------------------------------------------

def get_created_group(rank_list: Union[list[int], tuple[int, ...]]):
    """Return an existing process group by rank list, or ``None``."""
    group_key = str(tuple(sorted(rank_list)))
    if group_key in EXISTING_COMM_GROUPS:
        return EXISTING_COMM_GROUPS[group_key]
    return None


def split_group(parent_pg=None,
                split_ranks: Optional[list] = None,
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
    """Create split groups for every rank list in *split_ranks*.

    Returns the split process group relative to the current rank id.
    """
    del parent_pg, timeout, group_desc
    if split_ranks is None or len(split_ranks) == 0:
        raise ValueError("split_ranks cannot be None or empty")

    split_group_pg = None
    for split_rank in split_ranks:
        dist_group = get_created_group(split_rank)
        if dist_group is None:
            dist_group = dist.new_group(ranks=split_rank, pg_options=pg_options)
            EXISTING_COMM_GROUPS[str(tuple(sorted(split_rank)))] = dist_group
        if dist.get_rank() in split_rank:
            split_group_pg = dist_group

    return split_group_pg


def init_process_group(*args, **kwargs) -> None:
    """Initialize the default torch distributed process group."""
    if not dist.is_initialized():
        dist.init_process_group(*args, **kwargs)


# ---------------------------------------------------------------------------
# Differentiable collectives
163
164
165
166
167
168
169
170
171
    _OP_MAP['mean'] = dist.ReduceOp.AVG
else:
    # Fallback for older torch versions if necessary, though this might require manual division upstream
    # Assuming standard behavior where 'mean' implies native AVG support or upstream handling
    _OP_MAP['mean'] = dist.ReduceOp.SUM


def resolve_reduce_op(op: Union[str, Any]) -> Any:
    """Resolve a string op name (or pass through an already-resolved ``ReduceOp``)."""
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264

    @staticmethod
    def forward(ctx, tensor: Tensor, peer_rank: int, group) -> Tensor:  # pylint: disable=arguments-differ
        """Perform symmetric bidirectional P2P exchange with ``peer_rank``."""
        ctx.peer_rank = peer_rank
        ctx.group = group
        send_buf = tensor.contiguous()
        recv_buf = torch.empty_like(send_buf)
        reqs = dist.batch_isend_irecv([
            dist.P2POp(dist.isend, send_buf, peer_rank, group),
            dist.P2POp(dist.irecv, recv_buf, peer_rank, group),
        ])
        for req in reqs:
            req.wait()
        return recv_buf

    @staticmethod
    def backward(ctx, grad_output: Tensor):
        """Perform symmetric P2P exchange for the backward gradient pass."""
        send_buf = grad_output.contiguous()
        recv_buf = torch.empty_like(send_buf)
        reqs = dist.batch_isend_irecv([
            dist.P2POp(dist.isend, send_buf, ctx.peer_rank, ctx.group),
            dist.P2POp(dist.irecv, recv_buf, ctx.peer_rank, ctx.group),
        ])
        for req in reqs:
            req.wait()
        return recv_buf, None, None


class _TorchDifferentiableVariableAllGather(torch.autograd.Function):
    """Variable dim-zero all-gather with an uneven reduce-scatter backward."""
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285

    @staticmethod
    def forward(ctx, input_tensor, output_splits, group):  # pylint: disable=arguments-differ
        """Gather each rank's true row count without replicating inputs for A2A."""
        if input_tensor.ndim == 0:
            raise ValueError("variable all-gather input must have at least one dimension")
        splits = tuple(output_splits)
        if not splits:
            raise ValueError("output_splits must contain at least one group rank")
        if any(not isinstance(rows, int) or isinstance(rows, bool) or rows < 0 for rows in splits):
            raise ValueError(f"output_splits must contain non-negative integers, got {splits!r}")

        group_rank = dist.get_rank(group=group)
        if group_rank < 0 or group_rank >= len(splits):
            raise ValueError(f"group rank must be in [0, {len(splits)}), got {group_rank}")
        if input_tensor.shape[0] != splits[group_rank]:
            raise ValueError(
                "variable all-gather local rows must match output_splits at the group rank, "
                f"got local_rows={input_tensor.shape[0]}, group_rank={group_rank}, "
                f"output_splits={splits!r}"
            )
283
284
285
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
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
                f"got local_rows={input_tensor.shape[0]}, group_rank={group_rank}, "
                f"output_splits={splits!r}"
            )

        input_tensor = input_tensor.contiguous()
        feature_shape = tuple(input_tensor.shape[1:])
        if input_tensor.device.type == "npu":
            gathered = [input_tensor.new_empty((rows, *feature_shape)) for rows in splits]
            dist.all_gather(gathered, input_tensor, group=group)
        else:
            max_rows = max(splits)
            if max_rows == 0:
                gathered = [input_tensor.new_empty((0, *feature_shape)) for _ in splits]
            else:
                padded = input_tensor.new_zeros((max_rows, *feature_shape))
                if input_tensor.shape[0] > 0:
                    padded[:input_tensor.shape[0]].copy_(input_tensor)
                padded_outputs = [torch.empty_like(padded) for _ in splits]
                dist.all_gather(padded_outputs, padded, group=group)
                gathered = [
                    output[:rows].contiguous()
                    for output, rows in zip(padded_outputs, splits)
                ]

        ctx.output_splits = splits
        ctx.group = group
        ctx.group_rank = group_rank
        return torch.cat(gathered, dim=0)

    @staticmethod
    def backward(ctx, grad_output):
        """Sum replicated output gradients and return this rank's uneven shard."""
        output_rows = ctx.output_splits[ctx.group_rank]
        output = grad_output.new_empty((output_rows, *grad_output.shape[1:]))
        if sum(ctx.output_splits) == 0:
            return output, None, None

        grad_output = grad_output.contiguous()
        if grad_output.device.type == "npu":
            from torch_npu.distributed import reduce_scatter_tensor_uneven  # pylint: disable=C0415
            reduce_scatter_tensor_uneven(
                output,
                grad_output,
                input_split_sizes=list(ctx.output_splits),
                op=dist.ReduceOp.SUM,
327
328
329
330
331
332
333
334
335
336
337
338
339
                op=dist.ReduceOp.SUM,
                group=ctx.group,
            )
        else:
            reduced = grad_output.clone()
            dist.all_reduce(reduced, op=dist.ReduceOp.SUM, group=ctx.group)
            start = sum(ctx.output_splits[:ctx.group_rank])
            output.copy_(reduced.narrow(0, start, output_rows))
        return output, None, None


def differentiable_all_gather_concat(data: Tensor, group, concat_size: int, concat_dim: int,
                                     rank_list=None) -> Tensor:
357
358
359
360
361
362
363
364
365
366
367
368
        _TorchContiguousGrad.apply(tensor)
        for tensor in dist_func.all_gather(data, group=group)
    ]
    if rank_list is not None:
        group_ranks = dist.get_process_group_ranks(group)
        if tuple(rank_list) != tuple(group_ranks):
            rank_to_idx = {int(rank): idx for idx, rank in enumerate(group_ranks)}
            output = [output[rank_to_idx[int(rank)]] for rank in rank_list]
    return torch.cat(output, dim=concat_dim)


def differentiable_all_to_all_single(input_tensor: Tensor, input_splits: Sequence[int],
404
405
406
407
408
409
410
411
412
413
414


def differentiable_all_to_all(input_data: Tensor, output_shape: Sequence[int], group) -> Tensor:
    """Autograd-aware all-to-all producing a tensor of ``output_shape``."""
    input_data = ensure_contiguous(input_data)
    output_tensor = torch.empty(output_shape, device=input_data.device, dtype=input_data.dtype)
    return dist_func.all_to_all_single(output_tensor, input_data, group=group)


def differentiable_all_reduce(data: Tensor, op: Union[str, Any], group) -> Tensor:
    """Autograd-aware all-reduce with string or ``ReduceOp`` *op*."""
428
429
430
431
432
433
434
435
436
    )

    # 'avg' maps to SUM in _OP_MAP, so the division stays manual.
    if op == 'avg':
        output_tensor = output_tensor / dev_num
    return output_tensor


def differentiable_variable_all_gather(
435
436
437
438
439
440
441
442

def differentiable_variable_all_gather(
        input_tensor: Tensor, output_splits: Sequence[int], group) -> Tensor:
    """Gather variable dim-zero shards on HCCL or Gloo with autograd support."""
    return _TorchDifferentiableVariableAllGather.apply(
        input_tensor, tuple(output_splits), group
    )

452
453
454
455
456
457
458
459
460
461
462


def p2p_exchange(tensor: Tensor, peer_rank: int, group=None) -> Tensor:
    """Symmetric bidirectional P2P exchange with ``peer_rank``."""
    if peer_rank == dist.get_rank(group):
        return tensor
    return _TorchP2PExchangeFunction.apply(tensor, peer_rank, group)


def exchange_splits_via_all_to_all(input_tensor: Tensor, group) -> Tensor:
    """All-to-all a per-rank split vector, returning the received counts.
hyper_parallel/distributed/_builder/tp_collective_lowering.py
40
41
42
43
44
45
46
47
48

    def execute(self, tensor: Any) -> Any:
        """Execute the differentiable collective selected during lowering."""
        if self.kind == "all_gather":
            return communication.differentiable_all_gather_concat(
                tensor,
                self.group,
                self.group_size,
                self.tensor_dim,
47
48
49
50
51
52
53
54
55
                self.group_size,
                self.tensor_dim,
            )
        if self.kind == "all_reduce":
            return communication.differentiable_all_reduce(
                tensor,
                self.reduce_op,
                self.group,
            )
53
54
55
56
57
58
59
60
61
                self.reduce_op,
                self.group,
            )
        if self.kind == "reduce_scatter":
            return communication.differentiable_reduce_scatter(
                tensor,
                self.group_size,
                self.tensor_dim,
                self.reduce_op,
61
62
63
64
65
66
67
68
69
                self.reduce_op,
                self.group,
            )
        if self.kind == "all_reduce_shard":
            reduced = communication.differentiable_all_reduce(
                tensor,
                self.reduce_op,
                self.group,
            )
hyper_parallel/distributed/context_parallel/attention.py
301
302
303
304
305
306
307
308
309
            f"multiple of 2 * cp_size ({2 * cp_mesh.size()})"
        )

    peer_rank = _head_tail_peer_rank(cp_mesh)
    query_peer = communication.p2p_exchange(
        query.narrow(2, local_q_len // 2, local_q_len // 2), peer_rank)
    global_key, global_value = flex_cp_allgather(key, value, 2, cp_mesh)
    keep_output = _run_head_tail_half(
        attention_fn, query.narrow(2, 0, local_q_len // 2),
hyper_parallel/distributed/context_parallel/collectives.py
466
467
468
469
470
471
472
473
        + list(range(scatter_dim))
        + list(range(scatter_dim + 1, split_ndim))
    )
    send = tensor.contiguous().reshape(split_shape).permute(permutation).contiguous()
    received = communication.differentiable_all_to_all(
        send, list(send.shape), cp_mesh.get_group())
    return _reconstruct_all_to_all(received, gather_dim)

hyper_parallel/distributed/context_parallel/gated_delta_net.py
160
161
162
163
164
165
166
167
168
    output_splits = [0] * cp_size
    if cp_rank > 0:
        output_splits[rank_to_group_index[rank_list[cp_rank - 1]]] = halo_width

    exchange_output = communication.differentiable_all_to_all_single(
        exchange_input,
        input_splits,
        output_splits,
        group=cp_group,
906
907
908
909
910
911
912
913
914
    a2a_input = a2a_input.contiguous()
    split_len = a2a_input.shape[0] // split_count
    input_splits = [split_len] * split_count
    output_splits = [split_len] * split_count
    output = communication.differentiable_all_to_all_single(
        a2a_input,
        input_splits,
        output_splits,
        group=device_mesh.get_group(),
hyper_parallel/distributed/context_parallel/kimi_delta_attention.py
189
190
191
192
193
194
195
196
197
    output_splits = [0] * cp_size
    if cp_rank > 0:
        output_splits[rank_to_group_index[rank_list[cp_rank - 1]]] = halo_width

    exchange_output = communication.differentiable_all_to_all_single(
        exchange_input,
        input_splits,
        output_splits,
        group=cp_group,
312
313
314
315
316
317
318
319
320
    all_to_all_input = all_to_all_input.flatten(0, 1)

    split_length = all_to_all_input.shape[0] // split_count
    splits = [split_length] * split_count
    output = communication.differentiable_all_to_all_single(
        all_to_all_input,
        splits,
        splits,
        group=device_mesh.get_group(),