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/data/batching/build_dataloader.py 87.5% 42,46
hyper_parallel/platform/torch/dry_run.py 53.8% 80-83,87-93,96-98,124-127,131,138-147,157-159,163,168-171,176,180,182-188,203,206,211,222,275-276,280-281,287,297-298,305,309,328-334,338,348,370-371,383,389,391,398-401,411-412,416-424,436-437,441-452,457-462,469-475,480-496,501-504,507-509,514-517,522-527,559,562,565,570,575,582,584,586,588,590,601,608,621-622,638-650,658,671,722,724,740,742,744,755,770,822,826,834,836,845-850,853-854,857,865-867,893-902,932-933,938,948,956,967-972,1013-1016,1020-1027,1053-1054,1072-1082,1087,1117-1120,1124-1131,1135-1137,1140-1145,1149,1151-1158,1161-1169,1174,1176-1181,1183-1188,1195-1205,1214-1215,1217,1222,1228-1230,1233-1235,1237,1248-1250,1254-1255,1257,1259,1262-1263,1266-1269,1272,1275-1283,1286,1288-1289,1291-1295,1297-1298,1301-1302,1304-1307,1311-1316,1324-1329,1333-1339,1343-1345,1353-1362,1364,1366,1368-1371,1373-1374,1376,1384,1392-1396,1405,1410-1412,1418,1424-1425,1428,1433,1442-1444,1447-1448,1452,1456,1461,1467,1469-1474,1478-1482,1485,1488,1491-1492,1495-1501,1513-1514,1516-1518,1521-1528,1537-1540,1544-1546,1551,1557-1559,1561,1576-1577,1580,1587-1589,1591-1592,1594,1596,1601-1603,1610-1617,1625-1626,1628-1631,1634-1636,1639-1641,1643,1646-1649,1652-1654,1660,1668,1671,1676-1677,1683-1684,1690,1697-1699,1711,1713-1714,1718-1719,1722-1723,1726-1728,1736,1738-1744,1750-1752,1760,1762-1764,1767-1768,1773-1774,1779,1782,1790,1792,1795-1796,1801-1803,1807-1811,1815-1816,1818-1819,1821-1828,1832,1836-1837,1839-1840,1843-1844,1847,1854,1856-1858,1861,1864-1865,1868-1869,1872-1874,1883,1885-1890,1893,1902-1906,1910-1913,1917-1925,1927,1949-1950,1978,2082-2083,2107,2115-2119,2126-2127,2166,2180-2184,2212,2224-2225,2228,2232-2237,2243,2247,2249,2266,2279-2284,2286-2290,2311-2316,2331,2362,2367,2398-2399,2413,2459-2461,2477-2482,2514,2540,2550,2621-2624,2716,2724-2730,2732-2735,2737-2739,2746-2747,2819,2822,2833,2873,2912-2914,2919-2928
hyper_parallel/platform/torch/dtensor.py 29.2% 28-31,33-36,57-58,63,89,99,140,149-150,154,160,165-167,184-189,191,203-206,213-217,360-362,370-371,375-376,487-488
hyper_parallel/platform/torch/fully_shard/param.py 33.3% 53-60,62-67,69-72,74-76,775,941,1048
hyper_parallel/platform/torch/fully_shard/param_group.py 50.0% 146-147
hyper_parallel/platform/torch/fully_shard/scheduler.py 0.0% 121-123,134
hyper_parallel/platform/torch/memory_report.py 29.3% 65-73,78-82,95-98,103-113
hyper_parallel/trainer/config/__init__.py 100%  
hyper_parallel/trainer/config/dry_run.py 100%  
hyper_parallel/trainer/config/parallelism.py 100%  
hyper_parallel/trainer/config/resolver.py 0.0% 265-268
hyper_parallel/trainer/config/trainer.py 100%  
hyper_parallel/trainer/dry_run.py 48.0% 98-100,112-113,115-117,142,150,158,162,164,167,173,175,179,203,207,210,214,226-228,251,257,264,276,285-288,293-311,319-325,337-341,346-348,353,363,368-369,373-380,384,388-390,394-395,400-403,408-410,415-416,418,420,434-448,452,461-464,466,476,481,485,496-498,500-502,504,506-507,509-512,514-515,519-521,546,558-572,576-577,581-584,588,593,597-604,606-607,618,621-626,632-635,642,664,815,821,840,842,857,859,864-869,878,886-890,892,927,1026,1028,1030-1039,1044-1048,1050,1064-1076,1082-1084,1086-1088,1100-1105,1110,1115,1117-1124,1126-1128,1138-1141,1144-1146,1150-1153,1155-1160,1164-1165,1174-1182,1186-1187,1191-1193,1198-1203
hyper_parallel/trainer/dry_run_data.py 79.1% 59,63-64,66,75-79
hyper_parallel/trainer/dry_run_pipeline.py 50.8% 54,62,67,96,101,104-105,110,131,136,140,142,175,177,182,184,186-187,227,235-237,241-246,250-252,256-257,261-262,266-270,274-275,279,284,288,292,296,300,309-312,316-319,323-325
hyper_parallel/trainer/dry_run_pipeline_assembly.py 34.0% 88,94,99-102,107-119,124-127,132-150,159-166,178-182,187-189,194,201,210-220,232,235-245,252,266,287,292,297,304,318-321,330,333,338-344,355,357,359,363,367,372,390,405-414,417-422,427,432-435,440-445,449-453,459-464,476-478,481-484,491
hyper_parallel/data/batching/build_dataloader.py
38
39
40
41
42
43
44
45
46
47
48
49
    from torchdata.stateful_dataloader import StatefulDataLoader as _TorchStatefulDataLoader
except ModuleNotFoundError as error:
    missing_module = str(error.name or "")
    if missing_module != "torchdata" and not missing_module.startswith("torchdata."):
        raise
    _TORCHDATA_IMPORT_ERROR = error
    _StatefulDataLoaderBase = _UnavailableStatefulDataLoader
else:
    _StatefulDataLoaderBase = _TorchStatefulDataLoader

logger = get_dataset_logger(__name__)

hyper_parallel/platform/torch/dry_run.py
 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
        Raises:
            RuntimeError: If the process was not launched by ``torchrun``.
            ValueError: If a launcher variable is outside its valid range.
        """
        required_names = ("RANK", "WORLD_SIZE", "LOCAL_RANK")
        missing_names = [name for name in required_names if name not in os.environ]
        if missing_names:
            raise RuntimeError(
                "Multi-rank HyperDryRun must be launched with torchrun; "
                f"missing environment variables: {missing_names}"
            )
        rank = int(os.environ["RANK"])
        world_size = int(os.environ["WORLD_SIZE"])
        local_rank = int(os.environ["LOCAL_RANK"])
        if world_size < 1:
            raise ValueError(f"WORLD_SIZE must be >= 1, got {world_size}")
        if rank < 0 or rank >= world_size:
            raise ValueError(
                f"RANK must be in [0, {world_size}), got {rank}"
            )
        if local_rank < 0:
            raise ValueError(f"LOCAL_RANK must be >= 0, got {local_rank}")
        return cls(rank=rank, world_size=world_size, local_rank=local_rank)


class DryRunBatchMocker:
    """Erase one CPU training batch while preserving its complete structure."""
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
        if isinstance(value, Mapping):
            return self._mock_mapping(value)
        if isinstance(value, (list, tuple)):
            return self._mock_sequence(value)
        if is_dataclass(value) and not isinstance(value, type):
            return self._mock_dataclass(value)
        if hasattr(value, "__dict__") and self._contains_tensor(vars(value)):
            raise ValueError(
                "Dry-run cannot mock tensor-bearing batch carrier "
                f"{type(value).__module__}.{type(value).__qualname__}"
            )
        return value

    def _mock_mapping(self, value: Mapping[Any, Any]) -> Mapping[Any, Any]:
        """Mock mapping values while preserving the mapping carrier type."""
        mocked_items = [(key, self.mock(item)) for key, item in value.items()]
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
        """Mock mapping values while preserving the mapping carrier type."""
        mocked_items = [(key, self.mock(item)) for key, item in value.items()]
        if type(value) is dict:  # pylint: disable=unidiomatic-typecheck
            return dict(mocked_items)
        try:
            mocked_mapping = copy.copy(value)
            mocked_mapping.clear()
            mocked_mapping.update(mocked_items)
            return mocked_mapping
        except (AttributeError, TypeError):
            try:
                return type(value)(mocked_items)
            except TypeError as exc:
                raise ValueError(
                    "Dry-run cannot preserve batch mapping carrier "
                    f"{type(value).__module__}.{type(value).__qualname__}"
                ) from exc
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
        """Mock list and tuple values while preserving named tuples."""
        mocked = tuple(self.mock(item) for item in value)
        if isinstance(value, list):
            return list(mocked)
        if hasattr(value, "_fields"):
            return type(value)(*mocked)
        return mocked

    def _mock_dataclass(self, value: Any) -> Any:
        """Mock initialized and deferred dataclass fields."""
        updates = {
            field.name: self.mock(getattr(value, field.name))
            for field in fields(value)
            if field.init
        }
        mocked_dataclass = replace(value, **updates)
        for field in fields(value):
            if not field.init:
                object.__setattr__(
                    mocked_dataclass,
                    field.name,
                    self.mock(getattr(value, field.name)),
                )
        return mocked_dataclass

    def _contains_tensor(self, value: Any) -> bool:
        """Return whether an unsupported object recursively owns a tensor."""
        import torch  # pylint: disable=C0415

        if isinstance(value, torch.Tensor):
            return True
        if isinstance(value, Mapping):
            return any(self._contains_tensor(item) for item in value.values())
        if isinstance(value, (list, tuple)):
            return any(self._contains_tensor(item) for item in value)
        return False


def derive_tp_target_counts(
        loss_inputs: Mapping[str, Any],
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
    labels = loss_inputs.get("shift_labels")
    if labels is None:
        labels = loss_inputs.get("labels")
        if labels is None:
            raise ValueError("Dry-run loss_inputs must contain labels")
        labels = labels[..., 1:]
    if not isinstance(labels, torch.Tensor):
        raise ValueError("Dry-run labels must be a torch.Tensor")
    valid_mask = labels.ne(IGNORE_INDEX)
    loss_mask = loss_inputs.get("loss_mask")
    if isinstance(loss_mask, torch.Tensor):
        if loss_mask.shape != labels.shape:
            loss_mask = loss_mask[..., -labels.shape[-1]:]
        valid_mask = valid_mask & loss_mask.reshape_as(labels).bool()
    valid_labels = labels[valid_mask]
    chunk_size = (vocab_size + tp_size - 1) // tp_size
    owned = tuple(
218
219
220
221
222
223
224
225
226
        for rank in range(tp_size)
    )
    valid_count = int(valid_labels.numel())
    if sum(owned) != valid_count:
        raise ValueError("Dry-run labels contain token IDs outside the configured vocabulary")
    return valid_count, owned


@dataclass(frozen=True)
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
    paths: tuple[str, ...] = ()

    def discover(self, model: Any, context: Any) -> list[str]:
        """Return logical targets supported by this handler."""
        del model, context
        return []

    def recognizes(self, target: str, module: Any) -> bool:
        """Return whether this handler can explain an unconfigured target."""
        del target, module
        return False

    def validate(self, target: str, path: str, inputs: Dict[str, Any], context: Any) -> None:
        """Validate one expanded rule."""
        del target, inputs, context
283
284
285
286
287
288
289
290
    def validate(self, target: str, path: str, inputs: Dict[str, Any], context: Any) -> None:
        """Validate one expanded rule."""
        del target, inputs, context
        if self.paths and path not in self.paths:
            raise ValueError(
                f"value dependency handler {self.name!r} path must be one of "
                f"{self.paths}, got {path!r}"
            )
293
294
295
296
297
298
299
300
301
302
            self, stage: ValueDependencyStage, target: str, path: str,
            inputs: Dict[str, Any], context: Any,
    ) -> ContextManager[Any]:
        """Return scoped runtime state for one expanded rule."""
        del stage, target, path, inputs, context
        return nullcontext()

    def prepare_batch(
            self, target: str, path: str, inputs: Dict[str, Any],
            batch: Dict[str, Any], context: Any,
301
302
303
304
305
306
307
308
309
310
311
312
            self, target: str, path: str, inputs: Dict[str, Any],
            batch: Dict[str, Any], context: Any,
    ) -> None:
        """Apply shape-only batch changes for one expanded rule."""
        del target, path, inputs, batch, context

    def resolved_metadata(self) -> Dict[str, Any]:
        """Return handler-specific resolved report metadata."""
        return {}


_VALUE_DEPENDENCY_HANDLER_FACTORIES: Dict[str, Callable[[], ValueDependencyHandler]] = {}
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
    Raises:
        ValueError: If the name is empty, already registered, or inconsistent
            with the constructed handler.
    """
    if not isinstance(name, str) or not name.strip():
        raise ValueError("value dependency handler name must be a non-empty string")
    if name in _VALUE_DEPENDENCY_HANDLER_FACTORIES:
        raise ValueError(f"value dependency handler {name!r} is already registered")
    handler = factory()
    if not isinstance(handler, ValueDependencyHandler) or handler.name != name:
        raise ValueError(
            f"value dependency handler factory for {name!r} must return "
            f"ValueDependencyHandler(name={name!r})"
        )
    _VALUE_DEPENDENCY_HANDLER_FACTORIES[name] = factory


def _is_supported_moe_module(module: Any) -> bool:
    """Return whether a module exposes a supported routed-expert layout."""
344
345
346
347
348
349
350
351
352
    if experts is None or (
            getattr(module, "gate", None) is None
            and getattr(module, "router", None) is None
    ):
        return False
    layouts = (
        ("w1", "w2", "w3"),
        ("gate_up_proj", "down_proj"),
        ("gate_proj", "up_proj", "down_proj"),
366
367
368
369
370
371
372
373
374
375
        return [name for name, module in model.named_modules() if name and _is_supported_moe_module(module)]

    def recognizes(self, target: str, module: Any) -> bool:
        """Recognize supported routed-expert modules."""
        del target
        return _is_supported_moe_module(module)

    def validate(self, target: str, path: str, inputs: Dict[str, Any], context: Any) -> None:
        """Validate path-specific MoE rule inputs."""
        super().validate(target, path, inputs, context)
379
380
381
382
383
384
385
386
387
            "explicit": {"source_expert_loads"},
        }[path]
        unknown = sorted(set(inputs) - allowed)
        if unknown:
            raise ValueError(f"MoE rule for {target!r} contains unsupported inputs: {unknown}")
        if path == "hotspot" and (
                not isinstance(inputs.get("hotspot_expert"), int)
                or isinstance(inputs.get("hotspot_expert"), bool)
                or inputs["hotspot_expert"] < 0
385
386
387
388
389
390
391
392
393
394
395
                not isinstance(inputs.get("hotspot_expert"), int)
                or isinstance(inputs.get("hotspot_expert"), bool)
                or inputs["hotspot_expert"] < 0
        ):
            raise ValueError(f"MoE hotspot rule for {target!r} requires non-negative hotspot_expert")
        if path == "explicit" and not isinstance(inputs.get("source_expert_loads"), list):
            raise ValueError(f"MoE explicit rule for {target!r} requires source_expert_loads")

    def install(
            self, stage: ValueDependencyStage, target: str, path: str,
            inputs: Dict[str, Any], context: Any,
394
395
396
397
398
399
400
401
402
403
404
405
            self, stage: ValueDependencyStage, target: str, path: str,
            inputs: Dict[str, Any], context: Any,
    ) -> ContextManager[Any]:
        """Install routed-expert replay once for the configured fake step."""
        del target, path, inputs
        if stage is ValueDependencyStage.FAKE_STEP:
            return _MoERoutingValueDependencyRuntime(context.profile, context.base)
        return nullcontext()


class BranchValueDependencyHandler(ValueDependencyHandler):
    """Provide explicit values to instrumented Python branch decisions."""
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
    name = "branch"

    def discover(self, model: Any, context: Any) -> list[str]:
        """Expose module FQNs and logical execution regions."""
        del context
        return [name for name, _ in model.named_modules() if name] + ["loss", "backward", "optimizer"]

    def validate(self, target: str, path: str, inputs: Dict[str, Any], context: Any) -> None:
        """Validate named scalar branch decisions."""
        del path, context
        decisions = inputs.get("decisions")
        if set(inputs) != {"decisions"} or not isinstance(decisions, dict) or not decisions:
            raise ValueError(f"branch rule for {target!r} requires a non-empty decisions mapping")
        for name, value in decisions.items():
            if not isinstance(name, str) or not name.strip():
                raise ValueError(f"branch rule for {target!r} contains an invalid decision name")
            if not isinstance(value, (bool, int, float, str)) and value is not None:
                raise ValueError(
                    f"branch decision {name!r} for {target!r} must be a YAML scalar"
                )

432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
    name = "operator_debug"

    def discover(self, model: Any, context: Any) -> list[str]:
        """Expose module FQNs and logical execution regions."""
        del context
        return [name for name, _ in model.named_modules() if name] + ["loss", "backward", "optimizer"]

    def validate(self, target: str, path: str, inputs: Dict[str, Any], context: Any) -> None:
        """Validate source anchors and mock return envelopes."""
        del path, context
        mocks = inputs.get("mocks")
        if set(inputs) != {"mocks"} or not isinstance(mocks, list) or not mocks:
            raise ValueError(f"operator_debug rule for {target!r} requires a non-empty mocks list")
        selectors = set()
        for index, mock in enumerate(mocks):
            location = f"operator_debug mock {index} for {target!r}"
            selector = _operator_debug_selector(mock, location)
            if selector in selectors:
                raise ValueError(f"{location} duplicates an earlier selector")
            selectors.add(selector)
            _validate_operator_return(mock["return"], location)


def _operator_debug_selector(mock: Any, location: str) -> tuple[Any, ...]:
    """Validate and return the stable selector for one operator mock."""
    if not isinstance(mock, dict) or set(mock) != {"source", "op", "occurrence", "return"}:
        raise ValueError(f"{location} must contain source, op, occurrence, and return")
    source = mock["source"]
    if not isinstance(source, dict) or set(source) != {"file", "function", "line"}:
        raise ValueError(f"{location}.source must contain file, function, and line")
    if (
            not isinstance(source["file"], str)
            or not isinstance(source["function"], str)
            or not isinstance(source["line"], int)
            or isinstance(source["line"], bool)
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
            or not isinstance(source["line"], int)
            or isinstance(source["line"], bool)
            or source["line"] < 1
    ):
        raise ValueError(f"{location} has an invalid source anchor")
    if not isinstance(mock["op"], str) or not mock["op"].startswith("aten."):
        raise ValueError(f"{location}.op must be an ATen overload name")
    occurrence = mock["occurrence"]
    if not isinstance(occurrence, int) or isinstance(occurrence, bool) or occurrence < 0:
        raise ValueError(f"{location}.occurrence must be a non-negative integer")
    return source["file"], source["function"], source["line"], mock["op"], occurrence


def _validate_operator_return(spec: Any, location: str) -> None:
    """Validate one recursive operator-debug return specification."""
    if not isinstance(spec, dict) or len(spec) != 1:
        raise ValueError(f"{location}.return must contain exactly one return kind")
    kind, value = next(iter(spec.items()))
    if kind == "scalar":
        if not isinstance(value, (bool, int, float)):
            raise ValueError(f"{location}.return.scalar must be bool, int, or float")
        return
    if kind == "tensor":
        _validate_operator_tensor_return(value, location)
        return
    if kind in ("tuple", "list"):
        _validate_operator_sequence_return(value, kind, location)
        return
    if kind == "by_global_rank":
        _validate_operator_rank_return(value, location)
        return
    raise ValueError(f"{location}.return uses unsupported kind {kind!r}")


def _validate_operator_tensor_return(value: Any, location: str) -> None:
    """Validate a tensor-shaped operator-debug return."""
    if not isinstance(value, dict) or set(value) != {"shape", "dtype"}:
        raise ValueError(f"{location}.return.tensor must contain shape and dtype")
    shape = value["shape"]
    if not isinstance(shape, list) or any(
            not isinstance(size, int) or isinstance(size, bool) or size < 0 for size in shape
    ):
        raise ValueError(f"{location}.return.tensor.shape must contain non-negative integers")
    if not isinstance(value["dtype"], str) or not value["dtype"]:
        raise ValueError(f"{location}.return.tensor.dtype must be a non-empty string")


def _validate_operator_sequence_return(value: Any, kind: str, location: str) -> None:
    """Validate a tuple- or list-shaped operator-debug return."""
    if not isinstance(value, list):
        raise ValueError(f"{location}.return.{kind} must be a list")
    for index, item in enumerate(value):
        _validate_operator_return(item, f"{location}.return.{kind}[{index}]")


def _validate_operator_rank_return(value: Any, location: str) -> None:
    """Validate a rank-indexed operator-debug return."""
    if not isinstance(value, dict) or not value:
        raise ValueError(f"{location}.return.by_global_rank must be a non-empty mapping")
    for rank, item in value.items():
        if not isinstance(rank, (int, str)) or not str(rank).isdigit():
            raise ValueError(f"{location}.return.by_global_rank keys must be non-negative ranks")
        _validate_operator_return(item, f"{location}.return.by_global_rank[{rank}]")


for _builtin_handler in (
        MoERoutingValueDependencyHandler,
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
    @staticmethod
    def _parse_rules(value_dependencies: Any) -> tuple[_ValueDependencyRule, ...]:
        """Validate the common rule envelope before model discovery."""
        if not isinstance(value_dependencies, dict):
            raise ValueError("dry_run.value_dependencies must be a mapping")
        unknown = sorted(set(value_dependencies) - {"rules"})
        if unknown:
            raise ValueError(f"dry_run.value_dependencies contains unknown fields: {unknown}")
        raw_rules = value_dependencies.get("rules", [])
        if not isinstance(raw_rules, list):
            raise ValueError("dry_run.value_dependencies.rules must be a list")
        rules = []
        for index, raw_rule in enumerate(raw_rules):
            location = f"dry_run.value_dependencies.rules[{index}]"
            if not isinstance(raw_rule, dict):
                raise ValueError(f"{location} must be a mapping")
            unknown_rule_fields = sorted(
                set(raw_rule) - {"match", "handler", "path", "inputs", "optional"}
            )
            if unknown_rule_fields:
                raise ValueError(f"{location} contains unknown fields: {unknown_rule_fields}")
            match = raw_rule.get("match")
            handler = raw_rule.get("handler")
            path = raw_rule.get("path")
            inputs = raw_rule.get("inputs", {})
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
            path = raw_rule.get("path")
            inputs = raw_rule.get("inputs", {})
            optional = raw_rule.get("optional", False)
            if not isinstance(match, str) or not match.strip():
                raise ValueError(f"{location}.match must be a non-empty string")
            if not isinstance(handler, str) or not handler.strip():
                raise ValueError(f"{location}.handler must be a non-empty string")
            if not isinstance(path, str) or not path.strip():
                raise ValueError(f"{location}.path must be a non-empty string")
            if not isinstance(inputs, dict):
                raise ValueError(f"{location}.inputs must be a mapping")
            if not isinstance(optional, bool):
                raise ValueError(f"{location}.optional must be a boolean")
            rules.append(_ValueDependencyRule(index, match, handler, path, inputs, optional))
        return tuple(rules)

    def bind_model(self, model: Any, context: Any = None) -> None:
597
598
599
600
601
602
603
604
605
        candidates: Dict[tuple[str, str], list[tuple[int, _ValueDependencyRule]]] = {}
        for rule in self._rules:
            handler = self._handlers.get(rule.handler)
            if handler is None:
                raise ValueError(
                    f"value dependency rule {rule.index} uses unknown handler {rule.handler!r}; "
                    f"registered handlers: {sorted(self._handlers)}"
                )
            supported_targets = handler.discover(model, context)
604
605
606
607
608
609
610
611
612
                )
            supported_targets = handler.discover(model, context)
            matches = [target for target in supported_targets if fnmatchcase(target, rule.match)]
            if not matches and not rule.optional:
                raise ValueError(
                    f"value dependency rule {rule.index} match {rule.match!r} did not match "
                    f"any target supported by handler {rule.handler!r}"
                )
            exact = int(rule.match in matches)
617
618
619
620
621
622
623
624
625
626
        for (handler_name, target), target_candidates in candidates.items():
            best_priority = max(priority for priority, _ in target_candidates)
            best_rules = [rule for priority, rule in target_candidates if priority == best_priority]
            if len(best_rules) > 1:
                indexes = [rule.index for rule in best_rules]
                raise ValueError(
                    f"value dependency rules {indexes} conflict for target {target!r} "
                    f"and handler {handler_name!r}"
                )
            rule = best_rules[0]
634
635
636
637
638
639
640
641
642
643
644
645
646
647
648
649
650
651
652
653
654
        ``bind_model`` must run on the complete model before this projection.
        Rules are intentionally not glob-expanded again: a valid global target
        may be owned by a different pipeline rank.
        """
        self._target_modules = dict(model.named_modules())
        projected = {}
        for handler_name, targets in self._resolved.items():
            handler = self._handlers[handler_name]
            supported_targets = set(handler.discover(model, context))
            local_targets = {}
            for target, rule in targets.items():
                for local_target in supported_targets:
                    if local_target == target or local_target.endswith(f".{target}"):
                        local_targets[local_target] = rule
            if local_targets:
                projected[handler_name] = local_targets
        self._resolved = projected

    def resolved_rules(self, handler: str) -> Dict[str, _ValueDependencyRule]:
        """Return expanded rules for one handler."""
        return dict(self._resolved.get(handler, {}))
654
655
656
657
658
659
660
661
662
        return dict(self._resolved.get(handler, {}))

    def handler(self, name: str) -> ValueDependencyHandler:
        """Return one fresh profile-owned handler instance."""
        return self._handlers[name]

    @property
    def moe_layers(self) -> dict[str, Any]:
        """Return per-layer routing overrides."""
667
668
669
670
671
672
673
674
675

    @property
    def moe_enabled(self) -> bool:
        """Return whether the MoE adapter is enabled."""
        return bool(self._resolved.get("moe_routing"))

    def metadata(self) -> Dict[str, Any]:
        """Return a JSON-safe profile summary for the memory report."""
        return {
718
719
720
721
722
723
724
725
726
727
728
            )

        hotspot = int(layer.get("hotspot_expert", 0))
        if policy == "hotspot" and not 0 <= hotspot < num_experts:
            raise ValueError(f"MoE hotspot_expert for {module_name!r} must be in [0, {num_experts})")
        if policy == "hotspot":
            row = tuple(expected if index == hotspot else 0 for index in range(num_experts))
        else:
            base, remainder = divmod(expected, num_experts)
            row = tuple(base + int(index < remainder) for index in range(num_experts))
        return tuple(row for _ in range(ep_size))
736
737
738
739
740
741
742
743
744
745
746
747
748
            local_expert_count: int,
    ) -> int:
        """Validate the expert layout and return this rank's first expert."""
        if local_expert_count < 1:
            raise ValueError(f"MoE {module_name!r} has no local experts")
        if local_expert_count == num_experts:
            return 0
        if num_experts % ep_size or local_expert_count != num_experts // ep_size:
            raise ValueError(f"MoE {module_name!r} local expert layout is incompatible with EP={ep_size}")
        return ep_rank * local_expert_count

    @staticmethod
    def _expert_ranges(
751
752
753
754
755
756
757
758
759
            local_expert_count: int,
    ) -> tuple[tuple[int, int], ...]:
        """Return the global expert range owned by each destination rank."""
        if local_expert_count == num_experts:
            return ((0, num_experts),)
        return tuple(
            (rank * local_expert_count, (rank + 1) * local_expert_count)
            for rank in range(ep_size)
        )
766
767
768
769
770
771
772
773
774
        num_experts = int(module.num_experts)
        expected = local_tokens * int(module.top_k)
        layer = self.moe_layers.get(module_name)
        if layer is None:
            raise ValueError(f"MoE module {module_name!r} has no value-dependency rule")
        matrix = self._moe_routing_matrix(layer, ep_size, num_experts, expected, module_name)
        start = self._local_expert_start(
            module_name,
            ep_size,
818
819
820
821
822
823
824
825
826
827
828
829
830
            expected: int, module_name: str,
    ) -> tuple[tuple[int, ...], ...]:
        """Validate the configured aggregate EP traffic matrix."""
        if not isinstance(source_loads, list) or len(source_loads) != ep_size:
            raise ValueError(f"MoE explicit source_expert_loads for {module_name!r} must have {ep_size} rows")
        matrix = []
        for source_rank, row in enumerate(source_loads):
            if not isinstance(row, list) or len(row) != num_experts:
                raise ValueError(
                    f"MoE explicit source_expert_loads row {source_rank} for {module_name!r} "
                    f"must have {num_experts} entries"
                )
            if any(
830
831
832
833
834
835
836
837
838
839
840
            if any(
                    not isinstance(count, int) or isinstance(count, bool) or count < 0
                    for count in row
            ):
                raise ValueError("MoE explicit source_expert_loads entries must be non-negative integers")
            if sum(row) != expected:
                raise ValueError(
                    f"MoE explicit source_expert_loads row {source_rank} for {module_name!r} must sum to "
                    f"local_tokens * top_k ({expected}), got {sum(row)}"
                )
            matrix.append(tuple(row))
841
842
843
844
845
846
847
848
849
850
851
852
853
854
855
856
857
858
859
860
861
        return tuple(matrix)

    def tp_cross_entropy_target_count(self, target_count: int, tp_size: int, tp_rank: int) -> tuple[int, int]:
        """Return CPU-derived valid and local-owned target counts for TP CE."""
        counts = self._tp_cross_entropy_counts
        if counts is None:
            raise ValueError("TP cross-entropy counts were not derived from the training batch")
        valid, owned = counts
        if len(owned) != tp_size or not 0 <= tp_rank < tp_size:
            raise ValueError(
                f"CPU-derived TP target counts contain {len(owned)} ranks, expected {tp_size}"
            )
        if valid > target_count:
            raise ValueError(
                f"CPU-derived valid token count {valid} exceeds FakeTensor target size {target_count}"
            )
        return valid, owned[tp_rank]

    def configure_tp_cross_entropy_counts(
            self,
            valid_token_count: int,
861
862
863
864
865
866
867
868
869
870
871
            valid_token_count: int,
            target_tokens_per_rank: tuple[int, ...],
    ) -> None:
        """Install token ownership computed from the real CPU training batch."""
        if valid_token_count < 0 or sum(target_tokens_per_rank) != valid_token_count:
            raise ValueError("CPU-derived TP target counts must be non-negative and sum to the valid count")
        self._tp_cross_entropy_counts = (valid_token_count, target_tokens_per_rank)


_T = TypeVar("_T")
_ACTIVE_VALUE_DEPENDENCY_MANAGER: ContextVar[Optional["ValueDependencyManager"]] = ContextVar(
889
890
891
892
893
894
895
896
897
898
899
900
901
902
903
904
905
906

    Raises:
        ValueError: If ``name`` is empty or ``default_factory`` is not callable.
    """
    if not isinstance(name, str) or not name.strip():
        raise ValueError("value dependency decision name must be a non-empty string")
    if not callable(default_factory):
        raise ValueError("value dependency decision default_factory must be callable")
    manager = _ACTIVE_VALUE_DEPENDENCY_MANAGER.get()
    if manager is not None:
        configured, value = manager.resolve_decision(name)
        if configured:
            return value
    return default_factory()


class UnconfiguredValueDependencyError(RuntimeError):
    """Actionable error for an unconfigured FakeTensor value dependency."""
928
929
930
931
932
933
934
935
936
937
938
939
940
941
942

    @contextmanager
    def parallelize_context(self) -> Iterator[None]:
        """Install value dependencies required while parallelization is applied."""
        with ExitStack() as stack:
            self._install_handler_contexts(
                stack,
                ValueDependencyStage.PARALLELIZE,
                SimpleNamespace(profile=self._profile, model=self._model, runtime=self._runtime),
            )
            yield

    def fake_step_context(self, base: Any) -> "ValueDependencyManager":
        """Prepare this manager to enter the complete fake-step scope."""
        self._base = base
944
945
946
947
948
949
950
951
952

    def __enter__(self) -> "ValueDependencyManager":
        """Install semantic overrides, module scopes, and dispatch interception."""
        if self._base is None:
            raise RuntimeError("ValueDependencyManager.fake_step_context must be configured before entry")
        self._install_handler_contexts(
            self._stack,
            ValueDependencyStage.FAKE_STEP,
            SimpleNamespace(profile=self._profile, base=self._base, runtime=self._runtime),
952
953
954
955
956
957
958
959
960
            SimpleNamespace(profile=self._profile, base=self._base, runtime=self._runtime),
        )
        mesh = getattr(self._base, "mesh", None)
        if bool(getattr(mesh, "loss_parallel", False)):
            self._stack.enter_context(_TPCrossEntropyValueDependencyRuntime(self._profile))
        self._install_module_hooks()
        self._active_token = _ACTIVE_VALUE_DEPENDENCY_MANAGER.set(self)
        self._stack.enter_context(self._build_operator_debug_mode())
        return self
963
964
965
966
967
968
969
970
971
972
973
974
975
976
            self, stack: ExitStack, stage: ValueDependencyStage, context: Any,
    ) -> None:
        """Install semantic handlers once and external handlers per expanded target."""
        for handler_name, rules in self._profile._resolved.items():
            if handler_name in ("branch", "operator_debug") or not rules:
                continue
            handler = self._profile.handler(handler_name)
            selected_rules = [next(iter(rules.items()))] if handler_name in self._SEMANTIC_HANDLERS else rules.items()
            for target, rule in selected_rules:
                stack.enter_context(handler.install(stage, target, rule.path, rule.inputs, context))

    def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> bool:
        """Restore every scope and reject unused configured decisions/mocks."""
        del traceback
1009
1010
1011
1012
1013
1014
1015
1016
1017
1018
1019
1020
1021
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
            _VALUE_DEPENDENCY_TARGET_STACK.reset(token)

    def prepare_batch(self, batch: Dict[str, Any], context: Any) -> None:
        """Run configured handler batch preparation hooks."""
        for handler_name, rules in self._profile._resolved.items():
            handler = self._profile.handler(handler_name)
            for target, rule in rules.items():
                handler.prepare_batch(target, rule.path, rule.inputs, batch, context)

    def resolve_decision(self, name: str) -> tuple[bool, Any]:
        """Resolve a configured branch decision for the active innermost target."""
        target = self._current_target()
        if target is None:
            return False, None
        rule = self._profile.resolved_rules("branch").get(target)
        if rule is None or name not in rule.inputs["decisions"]:
            return False, None
        self._consumed_decisions.add((target, name))
        return True, rule.inputs["decisions"][name]

    def _install_module_hooks(self) -> None:
        """Track the innermost executing module with exception-safe hooks."""
        for module_name, module in self._model.named_modules():
1049
1050
1051
1052
1053
1054
1055
1056
1057
1058

            self._hook_handles.append(module.register_forward_pre_hook(_pre_hook))
            try:
                handle = module.register_forward_hook(_post_hook, always_call=True)
            except TypeError:
                handle = module.register_forward_hook(_post_hook)
            self._hook_handles.append(handle)

    def _clear_target_operator_counts(self, target: str) -> None:
        """Reset occurrence counters for one new target invocation."""
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091

    @staticmethod
    def _source_anchor() -> Dict[str, Any]:
        """Resolve the first user-code frame for the current ATen dispatch."""
        this_file = os.path.abspath(__file__)
        for frame_info in inspect.stack()[2:]:
            filename = os.path.abspath(frame_info.filename)
            if filename == this_file or f"{os.sep}site-packages{os.sep}torch{os.sep}" in filename:
                continue
            owner = frame_info.frame.f_locals.get("self")
            function = frame_info.function
            if owner is not None:
                function = f"{type(owner).__qualname__}.{function}"
            return {"file": frame_info.filename, "function": function, "line": frame_info.lineno}
        return {"file": "<unknown>", "function": "<unknown>", "line": 0}

    @staticmethod
    def _source_matches(configured: Dict[str, Any], actual: Dict[str, Any]) -> bool:
        """Return whether an actual frame matches a strict configured anchor."""
        return (
            actual["file"].endswith(configured["file"])
            and actual["function"] == configured["function"]
            and actual["line"] == configured["line"]
        )
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
1133
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
        target = self._current_target()
        op_name = str(func)
        rule = self._profile.resolved_rules("operator_debug").get(target or "")
        if rule is not None:
            matching_op_mocks = [mock for mock in rule.inputs["mocks"] if mock["op"] == op_name]
            if matching_op_mocks:
                source = self._source_anchor()
                anchored = [
                    mock for mock in matching_op_mocks
                    if self._source_matches(mock["source"], source)
                ]
                if anchored:
                    source_key = (source["file"], source["function"], source["line"])
                    count_key = (target, *source_key, op_name)
                    occurrence = self._operator_counts.get(count_key, 0)
                    self._operator_counts[count_key] = occurrence + 1
                    selected = [mock for mock in anchored if mock["occurrence"] == occurrence]
                    if len(selected) != 1:
                        raise ValueError(
                            f"operator_debug target {target!r} op {op_name} source {source} "
                            f"has no unique mock for occurrence {occurrence}"
                        )
                    mock = selected[0]
                    self._consumed_operator_mocks.add((target, id(mock)))
                    return self._build_operator_return(mock["return"], args, kwargs)
        try:
            return func(*args, **kwargs)
        except (DataDependentOutputException, DynamicOutputShapeException) as error:
            source = self._source_anchor()
            count_key = (target, source["file"], source["function"], source["line"], op_name)
            occurrence = self._operator_counts.get(count_key, 0)
            self._operator_counts[count_key] = occurrence + 1
            raise self._unconfigured_error(target, op_name, args, source, occurrence) from error

    def _build_operator_return(self, spec: Dict[str, Any], args: Any, kwargs: Dict[str, Any]) -> Any:
        """Construct one recursive operator-debug return value."""
        import torch  # pylint: disable=C0415

        kind, value = next(iter(spec.items()))
        if kind == "scalar":
            return value
        if kind == "by_global_rank":
            rank_key = str(self._runtime.rank)
            selected = value.get(self._runtime.rank, value.get(rank_key))
            if selected is None:
                raise ValueError(
                    f"operator_debug return has no by_global_rank entry for rank {self._runtime.rank}"
                )
            return self._build_operator_return(selected, args, kwargs)
        if kind in ("tuple", "list"):
            built = [self._build_operator_return(item, args, kwargs) for item in value]
            return tuple(built) if kind == "tuple" else built
        dtype = getattr(torch, value["dtype"], None)
        if not isinstance(dtype, torch.dtype):
            raise ValueError(f"operator_debug uses unknown torch dtype {value['dtype']!r}")
        device = self._first_tensor_device((args, kwargs))
        return torch.empty(tuple(value["shape"]), dtype=dtype, device=device)

    @staticmethod
    def _first_tensor_device(value: Any) -> Any:
        """Return the first tensor device in a nested operator argument tree."""
        import torch  # pylint: disable=C0415

        if isinstance(value, torch.Tensor):
            return value.device
        if isinstance(value, dict):
            values = value.values()
        elif isinstance(value, (tuple, list)):
            values = value
        else:
            return torch.device("cpu")
        for item in values:
            device = ValueDependencyManager._first_tensor_device(item)
            if device.type != "cpu" or isinstance(item, torch.Tensor):
                return device
        return torch.device("cpu")

    def _unconfigured_error(
            self, target: Optional[str], op_name: str, args: Any,
            source: Dict[str, Any], occurrence: int,
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
1207
1208
1209
            self, target: Optional[str], op_name: str, args: Any,
            source: Dict[str, Any], occurrence: int,
    ) -> Exception:
        """Build an actionable failure with semantic handler suggestions."""
        module = self._module_by_name.get(target or "")
        available = []
        if target is not None:
            for name, handler in self._profile._handlers.items():
                if name not in ("branch", "operator_debug") and handler.recognizes(target, module):
                    available.append(f"  - {name}: {', '.join(handler.paths)}")
        tensor_metadata = []
        for arg in args:
            if hasattr(arg, "shape") and hasattr(arg, "dtype"):
                tensor_metadata.append(f"shape={tuple(arg.shape)}, dtype={arg.dtype}")
        lines = [
            "Unconfigured value dependency detected",
            f"Target: {target or '<outside configured scope>'}",
            f"Operator: {op_name}",
            f"Occurrence: {occurrence}",
1210
1211
1212
1213
1214
1215
1216
1217
1218
1219
1220
1221
1222
1223
1224
1225
1226
            f"Source: {source['file']}:{source['line']} ({source['function']})",
            f"Rank: {self._runtime.rank}",
            f"Inputs: {tensor_metadata}",
        ]
        if available:
            lines.extend(["Available handlers:", *available, "Add a dry_run.value_dependencies rule for this target."])
        else:
            lines.extend([
                "No semantic handler recognizes this dependency.",
                "Register a ValueDependencyHandler, instrument a branch decision, "
                "or use source-anchored operator_debug.",
            ])
        return UnconfiguredValueDependencyError("\n".join(lines))

    def _validate_consumption(self) -> None:
        """Fail closed when configured branch decisions or operator mocks were unused."""
        unused_decisions = []
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
    def _validate_consumption(self) -> None:
        """Fail closed when configured branch decisions or operator mocks were unused."""
        unused_decisions = []
        for target, rule in self._profile.resolved_rules("branch").items():
            for name in rule.inputs["decisions"]:
                if (target, name) not in self._consumed_decisions:
                    unused_decisions.append(f"{target}:{name}")
        unused_mocks = []
        for target, rule in self._profile.resolved_rules("operator_debug").items():
            for index, mock in enumerate(rule.inputs["mocks"]):
                if (target, id(mock)) not in self._consumed_operator_mocks:
                    unused_mocks.append(f"{target}:mock[{index}]")
        if unused_decisions or unused_mocks:
            raise ValueError(
                "unused value-dependency configuration: "
                f"decisions={unused_decisions}, operator_mocks={unused_mocks}"
            )
1244
1245
1246
1247
1248
1249
1250
1251
1252
1253
1254
1255
1256
1257
1258
1259
1260
1261
1262
1263
1264
1265
1266
1267
1268
1269
1270
1271
1272
1273
1274
1275
1276
1277
1278
1279
1280
1281
1282
1283
1284
1285
1286
1287
1288
1289
1290
1291
1292
1293
1294
1295
1296
1297
1298
1299
1300
1301
1302
1303
1304
1305
1306
1307
1308
1309
1310
1311
1312
1313
1314
1315
1316
1317
1318
1319
1320
    """Scoped TP cross-entropy replay owned by its semantic handler."""

    def __init__(self, profile: _DryRunValueProfile) -> None:
        """Store the resolved ownership profile."""
        self._profile = profile
        self._loss_parallel_ops = None
        self._original_cross_entropy_function = None

    def __enter__(self) -> "_TPCrossEntropyValueDependencyRuntime":
        """Patch the complete CE autograd Function for this dry-run scope."""
        from hyper_parallel.platform.torch import loss_parallel_ops  # pylint: disable=C0415
        import torch  # pylint: disable=C0415

        profile = self._profile

        class DryRunDistributedCrossEntropyFunction(torch.autograd.Function):
            """CE Function that preserves tensor paths without inspecting labels."""

            @staticmethod
            def forward(ctx: Any, input_local: Any, target: Any, weight: Any, ignore_index: int,
                        reduction: str, vocab_size: int, mesh: Any, mesh_dim: int) -> Any:
                """Run distributed log-softmax and allocate configured local NLL state."""
                del weight, ignore_index, vocab_size
                tp_size = int(mesh.size(mesh_dim))
                tp_rank = int(mesh.get_local_rank(mesh_dim))
                valid_count, local_count = profile.tp_cross_entropy_target_count(
                    int(target.numel()), tp_size, tp_rank,
                )
                log_probs = loss_parallel_ops.distributed_log_softmax(
                    input_local, dim=-1, mesh=mesh, mesh_dim=mesh_dim,
                )
                selected = torch.empty((local_count,), dtype=log_probs.dtype, device=log_probs.device)
                anchor = log_probs.sum() * 0.0 + selected.sum() * 0.0
                ctx.save_for_backward(log_probs)
                ctx.reduction = reduction
                ctx.valid_count = valid_count
                ctx.local_count = local_count
                if reduction == "none":
                    return torch.empty_like(target, dtype=log_probs.dtype) + anchor
                total = loss_parallel_ops.platform.differentiable_all_reduce(
                    anchor, op="sum", group=mesh.get_group(mesh_dim),
                )
                return total / max(valid_count, 1) if reduction == "mean" else total

            @staticmethod
            def backward(ctx: Any, grad_output: Any) -> tuple[Any, ...]:
                """Retain local softmax-gradient allocations without label indexing."""
                (log_probs,) = ctx.saved_tensors
                if ctx.reduction == "none":
                    grad_scale = grad_output.reshape(-1, 1)
                elif ctx.reduction == "mean":
                    grad_scale = grad_output / max(ctx.valid_count, 1)
                else:
                    grad_scale = grad_output
                selected_grad = torch.empty(
                    (ctx.local_count,), dtype=log_probs.dtype, device=log_probs.device,
                )
                grad_input = log_probs.exp() * grad_scale + selected_grad.sum() * 0.0
                return grad_input, None, None, None, None, None, None, None

        self._loss_parallel_ops = loss_parallel_ops
        self._original_cross_entropy_function = loss_parallel_ops.DistributedCrossEntropyFunction
        loss_parallel_ops.DistributedCrossEntropyFunction = DryRunDistributedCrossEntropyFunction
        return self

    def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> bool:
        """Restore the real TP cross-entropy autograd Function."""
        del exc_type, exc_value, traceback
        if self._loss_parallel_ops is not None:
            self._loss_parallel_ops.DistributedCrossEntropyFunction = self._original_cross_entropy_function
            self._loss_parallel_ops = None
            self._original_cross_entropy_function = None
        return False


class _MoERoutingValueDependencyRuntimeBase:
    """Scoped routed-expert replay used by the MoE semantic handler."""
1320
1321
1322
1323
1324
1325
1326
1327
1328
1329
1330
1331
1332
1333
1334
1335
1336
1337
1338
1339
1340
1341
1342
1343
1344
1345
1346
1347
1348
1349
    """Scoped routed-expert replay used by the MoE semantic handler."""

    def __init__(self, profile: _DryRunValueProfile, base: Any) -> None:
        """Create a scope for one prepared trainer."""
        self._profile = profile
        self._base = base
        self._patched_moe_modules = []
        self._moe_module_names: Dict[int, str] = {}
        self._ep_compute_module = None
        self._original_ep_compute = None

    def __enter__(self) -> "_MoERoutingValueDependencyRuntime":
        """Install configured MoE replay."""
        try:
            if self._profile.moe_enabled:
                self._patch_moe()
        except Exception:
            self._restore_moe()
            raise
        return self

    def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> bool:
        """Restore every patched symbol even if the simulated step failed."""
        del exc_type, exc_value, traceback
        self._restore_moe()
        return False


class _MoERoutingValueDependencyRuntime(_MoERoutingValueDependencyRuntimeBase):
    """Complete shape and memory replay for configured routed-expert modules."""
1349
1350
1351
1352
1353
1354
1355
1356
1357
1358
1359
1360
1361
1362
1363
1364
1365
1366
1367
1368
1369
1370
1371
1372
1373
1374
1375
1376
1377
1378
1379
1380
    """Complete shape and memory replay for configured routed-expert modules."""

    def _patch_moe(self) -> None:
        """Patch supported MoE execution modules without changing their classes."""
        module_names = set()
        configured_names = set(self._profile.moe_layers)
        for module_name, module in self._base.model.named_modules():
            if module_name in configured_names and self._is_supported_moe_module(module):
                module_names.add(module_name)
                self._moe_module_names[id(module)] = module_name
                if self._base.mesh.ep_size > 1:
                    continue
                had_forward = "forward" in module.__dict__
                original_forward = module.__dict__.get("forward")

                def dry_run_forward(current_module: Any, x: Any, name: str = module_name) -> Any:
                    """Execute the configured value-independent MoE path."""
                    return self._run_moe_forward(name, current_module, x)

                module.forward = MethodType(dry_run_forward, module)
                self._patched_moe_modules.append((module, had_forward, original_forward))
        if self._base.mesh.ep_size > 1 and module_names:
            from hyper_parallel.distributed.expert_parallel import recipes as ep_compute  # pylint: disable=C0415

            self._ep_compute_module = ep_compute
            self._original_ep_compute = ep_compute.ep_routed_forward

            def dry_run_ep_compute(
                    module: Any,
                    hidden_states: Any,
                    *,
                    router_fn: Any,
1380
1381
1382
1383
1384
1385
1386
1387
1388
                    router_fn: Any,
                    ep_group: Any,
            ) -> Any:
                """Execute configured EP shapes inside the existing local region."""
                return self._run_hf_native_ep_compute(
                    module,
                    hidden_states,
                    router_fn=router_fn,
                    ep_group=ep_group,
1388
1389
1390
1391
1392
1393
1394
1395
1396
1397
1398
1399
1400
                    ep_group=ep_group,
                    tp_group=None,
                )

            ep_compute.ep_routed_forward = dry_run_ep_compute
        unknown_names = configured_names - module_names
        if unknown_names:
            self._restore_moe()
            raise ValueError(
                "moe_routing rules contain unsupported module FQNs: "
                f"{sorted(unknown_names)}; supported MoE module FQNs discovered "
                f"in the parallelized model: {sorted(module_names)}"
            )
1401
1402
1403
1404
1405
1406
1407
1408
1409
1410
1411
1412
1413
1414
1415
1416

    @staticmethod
    def _is_supported_moe_module(module: Any) -> bool:
        """Return whether a module exposes one of the supported expert layouts."""
        return _is_supported_moe_module(module)

    @staticmethod
    def _moe_dimensions(module: Any) -> tuple[int, int]:
        """Return global expert count and top-k for a supported MoE module."""
        experts = module.experts
        config = getattr(module, "config", None)
        num_experts = (
            getattr(experts, "num_experts", None)
            or getattr(module, "num_experts", None)
            or getattr(config, "num_experts", None)
            or getattr(config, "n_routed_experts", None)
1414
1415
1416
1417
1418
1419
1420
1421
1422
            or getattr(module, "num_experts", None)
            or getattr(config, "num_experts", None)
            or getattr(config, "n_routed_experts", None)
        )
        top_k = (
            getattr(module, "top_k", None)
            or getattr(getattr(module, "gate", None), "top_k", None)
            or getattr(config, "num_experts_per_tok", None)
            or getattr(config, "top_k", None)
1420
1421
1422
1423
1424
1425
1426
1427
1428
1429
1430
1431
1432
1433
1434
1435
1436
            or getattr(getattr(module, "gate", None), "top_k", None)
            or getattr(config, "num_experts_per_tok", None)
            or getattr(config, "top_k", None)
        )
        if num_experts is None or top_k is None:
            raise ValueError(
                f"Cannot resolve expert count/top-k for {type(module).__name__}"
            )
        return int(num_experts), int(top_k)

    @staticmethod
    def _local_weight(weight: Any) -> Any:
        """Return a local expert-weight shard while preserving FakeTensor state."""
        return weight.to_local() if hasattr(weight, "to_local") else weight

    def _expert_weight_layout(self, module: Any) -> tuple[str, tuple[Any, ...]]:
        """Extract local expert weights from a supported MoE module.
1438
1439
1440
1441
1442
1443
1444
1445
1446
1447
1448
1449
1450
1451
1452
1453
1454
1455
1456
1457
1458
1459
1460
1461
1462
1463
1464
1465
        The framework ``MoE`` stores independent SwiGLU matrices ``w1/w2/w3``.
        Qwen3.5 packs the gate and up matrices together in ``gate_up_proj``.
        Both layouts retain the EP-sharded leading expert dimension.
        """
        experts = module.experts
        if all(hasattr(experts, name) for name in ("w1", "w2", "w3")):
            return "three_projection", tuple(
                self._local_weight(getattr(experts, name)) for name in ("w1", "w2", "w3")
            )
        if all(hasattr(experts, name) for name in ("gate_up_proj", "down_proj")):
            return "packed_gate_up", (
                self._local_weight(experts.gate_up_proj),
                self._local_weight(experts.down_proj),
            )
        if all(
                hasattr(experts, name)
                for name in ("gate_proj", "up_proj", "down_proj")
        ):
            return "three_projection", (
                self._local_weight(experts.gate_proj),
                self._local_weight(experts.down_proj),
                self._local_weight(experts.up_proj),
            )
        raise ValueError(
            f"MoE module {type(module).__name__!r} does not expose a supported expert-weight layout"
        )

    def _run_moe_forward(self, module_name: str, module: Any, x: Any) -> Any:
1463
1464
1465
1466
1467
1468
1469
1470
1471
1472
1473
1474
1475
1476
1477
1478
1479
1480
1481
1482
1483
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498
1499
1500
1501
1502
1503
1504
1505
        )

    def _run_moe_forward(self, module_name: str, module: Any, x: Any) -> Any:
        """Run a value-independent analogue of the real routed MoE path."""
        import torch  # pylint: disable=C0415

        batch_size, sequence_length, hidden_size = x.shape
        ep_size, ep_rank = self._ep_group_identity()
        layout, weights = self._expert_weight_layout(module)
        num_experts, top_k = self._moe_dimensions(module)
        routing_module = SimpleNamespace(num_experts=num_experts, top_k=top_k)
        plan = self._profile.moe_routing_plan(
            module_name, routing_module, ep_size, ep_rank,
            int(batch_size * sequence_length), int(weights[0].shape[0]),
        )
        x_flat = x.view(-1, hidden_size)
        shared_output = None
        if getattr(module, "shared_expert", None) is not None:
            shared_output = module.shared_expert(x_flat)
        routed_input, top_weights, permutation, inverse_permutation, token_counts = self._dry_run_route(
            module_name, module, x_flat, plan,
        )
        expert_input, dispatch_context = self._dry_run_dispatch(
            routed_input, token_counts, plan, ep_size, int(weights[0].shape[0]),
        )
        expert_output = self._run_grouped_experts(
            layout, weights, expert_input, plan.local_expert_loads,
        )
        combined = self._dry_run_combine(expert_output, dispatch_context)
        routed_output = self._dry_run_weight_and_unpermute(
            module, x_flat, combined, top_weights, permutation, inverse_permutation,
        )
        output = routed_output
        if shared_output is not None:
            shared_gate = torch.sigmoid(module.shared_expert_gate(x_flat))
            output = output + shared_gate * shared_output
        if hasattr(module, "last_aux_loss"):
            module.last_aux_loss = None
        return output.view(batch_size, sequence_length, hidden_size)

    def _run_hf_native_ep_compute(
            self,
            module: Any,
1509
1510
1511
1512
1513
1514
1515
1516
1517
1518
1519
1520
1521
1522
1523
1524
1525
1526
1527
1528
1529
1530
1531
1532
            ep_group: Any,
            tp_group: Any,
    ) -> Any:
        """Replay the new Trainer's EP lifecycle with configured split sizes."""
        import torch  # pylint: disable=C0415
        import torch.distributed as dist  # pylint: disable=C0415

        module_name = self._moe_module_names.get(id(module))
        if module_name is None:
            raise ValueError(
                f"Dry-run EP compute received an unregistered {type(module).__name__}"
            )
        ep_size = int(ep_group.size())
        ep_rank = int(dist.get_rank(group=ep_group))
        num_experts, top_k = self._moe_dimensions(module)
        local_expert_count = int(module.experts.local_expert_count)
        batch_size, sequence_length, hidden_size = hidden_states.shape
        local_tokens = int(batch_size * sequence_length)
        routing_module = SimpleNamespace(num_experts=num_experts, top_k=top_k)
        plan = self._profile.moe_routing_plan(
            module_name,
            routing_module,
            ep_size,
            ep_rank,
1533
1534
1535
1536
1537
1538
1539
1540
1541
1542
1543
1544
1545
1546
1547
1548
1549
1550
1551
1552
1553
1554
1555
            local_tokens,
            local_expert_count,
        )

        flattened_states = hidden_states.reshape(-1, hidden_size)
        topk_indices, topk_weights = router_fn(module, hidden_states)
        flattened_weights = topk_weights.reshape(-1).to(flattened_states.dtype)
        source_indices = torch.arange(
            local_tokens,
            device=flattened_states.device,
        ).repeat_interleave(top_k)
        dispatch_order = topk_indices.reshape(-1).argsort()
        dispatched_states = flattened_states[source_indices[dispatch_order]].contiguous()
        received_states = self._differentiable_all_to_all(
            dispatched_states,
            plan.input_split_sizes,
            plan.output_split_sizes,
        )
        received_expert_indices = torch.empty(
            (sum(plan.output_split_sizes),),
            dtype=torch.int64,
            device=hidden_states.device,
        )
1553
1554
1555
1556
1557
1558
1559
1560
1561
1562
1563
1564
1565
            dtype=torch.int64,
            device=hidden_states.device,
        )

        experts = module.experts
        had_forward = "forward" in experts.__dict__
        original_forward = experts.__dict__.get("forward")

        def grouped_forward(
                current_experts: Any,
                dispatched: Any,
                local_indices: Any,
        ) -> Any:
1572
1573
1574
1575
1576
1577
1578
1579
1580
1581
1582
1583
1584

            Returns:
                The grouped expert outputs.
            """
            del local_indices
            layout, weights = self._expert_weight_layout(
                SimpleNamespace(experts=current_experts)
            )
            return self._run_grouped_experts(
                layout,
                weights,
                dispatched,
                plan.local_expert_loads,
1583
1584
1585
1586
1587
1588
1589
1590
1591
1592
1593
1594
1595
1596
1597
1598
1599
1600
1601
1602
1603
1604
1605
1606
1607
                dispatched,
                plan.local_expert_loads,
            )

        experts.forward = MethodType(grouped_forward, experts)
        try:
            local_outputs = experts(received_states, received_expert_indices)
        finally:
            if had_forward:
                experts.forward = original_forward
            else:
                del experts.forward

        combined = self._differentiable_all_to_all(
            local_outputs.contiguous(),
            plan.output_split_sizes,
            plan.input_split_sizes,
        )
        flattened_outputs = torch.zeros_like(combined)
        flattened_outputs[dispatch_order] = combined
        output = (
            flattened_outputs * flattened_weights.unsqueeze(-1)
        ).view(local_tokens, top_k, hidden_size).sum(dim=1).view(
            batch_size,
            sequence_length,
1606
1607
1608
1609
1610
1611
1612
1613
1614
1615
1616
1617
1618
1619
1620
1621
            batch_size,
            sequence_length,
            hidden_size,
        )
        shared = getattr(module, "shared_experts", None)
        if shared is not None:
            if tp_group is None:
                raise ValueError("MoE shared_experts requires a TP group")
            shared_output = shared(hidden_states)
            dist.all_reduce(shared_output, group=tp_group)
            output = output + shared_output
        return output

    @staticmethod
    def _dry_run_route(
            module_name: str, module: Any, x_flat: Any,
1621
1622
1623
1624
1625
1626
1627
1628
1629
1630
1631
1632
1633
1634
1635
1636
1637
1638
1639
1640
1641
1642
1643
1644
1645
1646
1647
1648
1649
1650
1651
1652
1653
1654
1655
1656
1657
1658
            module_name: str, module: Any, x_flat: Any,
            plan: _DryRunMoERoutingPlan,
    ) -> tuple[Any, Any, Any, Any, Any]:
        """Replay value-independent router, sort, and histogram operators."""
        import torch  # pylint: disable=C0415
        from torch.nn import functional  # pylint: disable=C0415

        gate = getattr(module, "gate", None)
        gate_weight = getattr(gate, "weight", None)
        if gate_weight is None:
            raise ValueError(
                f"MoE module {module_name!r} must expose gate.weight for configured dry-run routing"
            )
        router_logits = x_flat @ gate_weight.transpose(0, 1)
        router_probs = functional.softmax(router_logits, dim=-1, dtype=torch.float32)
        top_weights, top_indices = torch.topk(
            router_probs, int(module.top_k), dim=-1,
        )
        top_weights = top_weights / top_weights.sum(dim=-1, keepdim=True)
        top_weights = top_weights.to(x_flat.dtype)
        module.router_logits = router_logits

        token_indices = torch.arange(
            x_flat.shape[0], device=x_flat.device,
        ).unsqueeze(1).expand(-1, int(module.top_k)).reshape(-1)
        expert_indices = top_indices.reshape(-1)
        permutation = torch.argsort(expert_indices, stable=True)
        inverse_permutation = torch.empty_like(permutation)
        inverse_permutation[permutation] = torch.arange(
            plan.outgoing_tokens, device=x_flat.device,
        )
        routed_input = x_flat[token_indices[permutation]]
        histc_input = expert_indices.float() if x_flat.device.type == "cpu" else expert_indices.int()
        token_counts = torch.histc(
            histc_input,
            bins=int(module.num_experts),
            min=0,
            max=int(module.num_experts) - 1,
1656
1657
1658
1659
1660
1661
1662
1663
1664
            bins=int(module.num_experts),
            min=0,
            max=int(module.num_experts) - 1,
        ).to(torch.int64)
        return routed_input, top_weights, permutation, inverse_permutation, token_counts

    def _dry_run_dispatch(
            self, routed_input: Any, token_counts: Any,
            plan: _DryRunMoERoutingPlan,
1664
1665
1666
1667
1668
1669
1670
1671
1672
1673
1674
1675
1676
1677
1678
1679
1680
1681
            plan: _DryRunMoERoutingPlan,
            ep_size: int, local_expert_count: int,
    ) -> tuple[Any, _DryRunMoEDispatchContext]:
        """Replay counts exchange, token dispatch, and expert permutation."""
        counts_output = self._dry_run_exchange_counts(
            token_counts, ep_size,
        )
        dispatched = self._differentiable_all_to_all(
            routed_input,
            plan.input_split_sizes,
            plan.output_split_sizes,
        )
        rank_major_shape = tuple(dispatched.shape)
        permutation = self._dry_run_rank_to_expert_indices(
            counts_output,
            sum(plan.output_split_sizes),
            ep_size,
            local_expert_count,
1679
1680
1681
1682
1683
1684
1685
1686
1687
1688
            sum(plan.output_split_sizes),
            ep_size,
            local_expert_count,
        )
        expert_input = dispatched[permutation]
        context = _DryRunMoEDispatchContext(
            rank_major_shape=rank_major_shape,
            permuted_indices=permutation,
            input_split_sizes=plan.input_split_sizes,
            output_split_sizes=plan.output_split_sizes,
1686
1687
1688
1689
1690
1691
1692
1693
1694
            permuted_indices=permutation,
            input_split_sizes=plan.input_split_sizes,
            output_split_sizes=plan.output_split_sizes,
        )
        return expert_input, context

    def _dry_run_combine(
            self, expert_output: Any,
            context: _DryRunMoEDispatchContext,
1693
1694
1695
1696
1697
1698
1699
1700
1701
1702
1703
            self, expert_output: Any,
            context: _DryRunMoEDispatchContext,
    ) -> Any:
        """Replay expert unpermutation and reverse token all-to-all."""
        rank_major_output = expert_output.new_zeros(*context.rank_major_shape)
        rank_major_output[context.permuted_indices] = expert_output
        return self._differentiable_all_to_all(
            rank_major_output,
            context.output_split_sizes,
            context.input_split_sizes,
        )
1707
1708
1709
1710
1711
1712
1713
1714
1715
1716
1717
1718
1719
1720
1721
1722
1723
1724
1725
1726
1727
1728
1729
1730
1731
1732
            module: Any, x_flat: Any, combined: Any, top_weights: Any,
            permutation: Any, inverse_permutation: Any,
    ) -> Any:
        """Apply routing weights and restore the original token order."""
        import torch  # pylint: disable=C0415

        sorted_weights = top_weights.reshape(-1)[permutation]
        use_fp32_combine = (
            getattr(module, "_hp_moe_tp_enabled", False)
            or getattr(module, "_hp_moe_ep_fp32_routing", False)
        )
        if use_fp32_combine:
            weighted = (
                combined.to(torch.float32) * sorted_weights.to(torch.float32).unsqueeze(-1)
            ).to(combined.dtype)
            unsorted = weighted[inverse_permutation]
            return unsorted.view(
                x_flat.shape[0], int(module.top_k), x_flat.shape[-1],
            ).sum(dim=1, dtype=torch.float32).to(x_flat.dtype)
        weighted = combined * sorted_weights.unsqueeze(-1)
        unsorted = weighted[inverse_permutation]
        return unsorted.view(
            x_flat.shape[0], int(module.top_k), x_flat.shape[-1],
        ).sum(dim=1).to(x_flat.dtype)

    def _dry_run_exchange_counts(
1732
1733
1734
1735
1736
1737
1738
1739
1740
1741
1742
1743
1744
1745
1746
1747
1748
    def _dry_run_exchange_counts(
            self, counts_input: Any, ep_size: int,
    ) -> Any:
        """Replay the non-differentiable count all-to-all at a fixed shape."""
        from hyper_parallel.platform import get_platform  # pylint: disable=C0415

        if ep_size == 1:
            return counts_input.clone()
        try:
            ep_group = self._ep_mesh().get_group()
        except (KeyError, TypeError, AttributeError) as error:
            raise ValueError("Configured MoE dry-run requires an accessible EP process group") from error
        counts_output, handle = get_platform().all_to_all_single(
            counts_input,
            output_shape=[counts_input.shape[0]],
            group=ep_group,
            async_op=True,
1746
1747
1748
1749
1750
1751
1752
1753
1754
1755
1756
            output_shape=[counts_input.shape[0]],
            group=ep_group,
            async_op=True,
        )
        if handle is not None:
            handle.wait()
        return counts_output

    @staticmethod
    def _dry_run_rank_to_expert_indices(
            counts: Any, total_tokens: int,
1756
1757
1758
1759
1760
1761
1762
1763
1764
1765
1766
1767
1768
1769
1770
1771
1772
1773
1774
1775
1776
1777
1778
1779
1780
1781
1782
1783
1784
1785
1786
            counts: Any, total_tokens: int,
            ep_size: int, local_expert_count: int,
    ) -> Any:
        """Replay rank-major to expert-major index construction at fixed shape."""
        import torch  # pylint: disable=C0415

        counts_2d = counts.view(ep_size, local_expert_count)
        source_offsets = counts.cumsum(0) - counts
        expert_major_offsets = source_offsets.view(
            ep_size, local_expert_count,
        ).transpose(0, 1).contiguous().view(-1)
        expert_major_counts = counts_2d.transpose(0, 1).contiguous().view(-1)
        block_source_starts = torch.repeat_interleave(
            expert_major_offsets,
            expert_major_counts,
            output_size=total_tokens,
        )
        destination_offsets = expert_major_counts.cumsum(0) - expert_major_counts
        destination_starts = torch.repeat_interleave(
            destination_offsets,
            expert_major_counts,
            output_size=total_tokens,
        )
        intra_block_offsets = torch.arange(
            total_tokens, device=counts.device,
        ) - destination_starts
        return (block_source_starts + intra_block_offsets).long()

    @staticmethod
    def _dry_run_grouped_matmul(
            input_tensor: Any, weight: Any,
1786
1787
1788
1789
1790
1791
1792
1793
1794
1795
1796
1797
1798
1799
1800
1801
1802
1803
1804
1805
1806
1807
1808
1809
1810
1811
1812
1813
1814
1815
1816
1817
1818
1819
1820
1821
1822
1823
1824
1825
1826
1827
1828
1829
1830
1831
1832
1833
1834
1835
1836
1837
1838
1839
1840
1841
1842
1843
1844
1845
1846
1847
1848
1849
1850
1851
            input_tensor: Any, weight: Any,
            expert_loads: tuple[int, ...],
    ) -> Any:
        """Replay grouped matmul events with one output and one weight gradient."""
        import torch  # pylint: disable=C0415

        class _DryRunGroupedMatmul(torch.autograd.Function):
            """Grouped matmul analogue with preallocated forward/backward buffers."""

            @staticmethod
            def forward(
                    ctx: Any, value: Any, grouped_weight: Any,
                    loads: tuple[int, ...],
            ) -> Any:
                """Compute packed expert outputs into one logical allocation."""
                ctx.save_for_backward(value, grouped_weight)
                ctx.expert_loads = loads
                output = torch.zeros(
                    value.shape[0], grouped_weight.shape[-1],
                    dtype=value.dtype, device=value.device,
                )
                start = 0
                for expert_index, load in enumerate(loads):
                    end = start + load
                    if load:
                        torch.mm(
                            value[start:end], grouped_weight[expert_index],
                            out=output[start:end],
                        )
                    start = end
                return output

            @staticmethod
            def backward(ctx: Any, grad_output: Any) -> tuple[Any, Any, None]:
                """Compute packed input and expert-weight gradients."""
                value, grouped_weight = ctx.saved_tensors
                grad_input = torch.zeros_like(value)
                grad_weight = torch.zeros_like(grouped_weight)
                start = 0
                for expert_index, load in enumerate(ctx.expert_loads):
                    end = start + load
                    if load:
                        torch.mm(
                            grad_output[start:end], grouped_weight[expert_index].transpose(0, 1),
                            out=grad_input[start:end],
                        )
                        torch.mm(
                            value[start:end].transpose(0, 1), grad_output[start:end],
                            out=grad_weight[expert_index],
                        )
                    start = end
                return grad_input, grad_weight, None

        if len(expert_loads) != int(weight.shape[0]):
            raise ValueError(
                "MoE dry-run expert load count must match the local expert weight count"
            )
        if sum(expert_loads) != int(input_tensor.shape[0]):
            raise ValueError(
                "MoE dry-run local expert loads must sum to the dispatched token count"
            )
        return _DryRunGroupedMatmul.apply(input_tensor, weight, expert_loads)

    def _run_grouped_experts(
            self, layout: str, weights: tuple[Any, ...], expert_input: Any,
            expert_loads: tuple[int, ...],
1850
1851
1852
1853
1854
1855
1856
1857
1858
1859
1860
1861
1862
1863
1864
1865
1866
1867
1868
1869
1870
1871
1872
1873
1874
1875
1876
1877
1878
            self, layout: str, weights: tuple[Any, ...], expert_input: Any,
            expert_loads: tuple[int, ...],
    ) -> Any:
        """Execute the supported expert layout without per-expert slicing."""
        from torch.nn import functional  # pylint: disable=C0415

        if layout == "three_projection":
            w1, w2, w3 = weights
            gate = self._dry_run_grouped_matmul(
                expert_input, w1.transpose(-2, -1), expert_loads,
            )
            up = self._dry_run_grouped_matmul(
                expert_input, w3.transpose(-2, -1), expert_loads,
            )
            hidden = functional.silu(gate) * up
            return self._dry_run_grouped_matmul(
                hidden, w2.transpose(-2, -1), expert_loads,
            )
        gate_up, down = weights
        gate_up_output = self._dry_run_grouped_matmul(
            expert_input, gate_up.transpose(-2, -1), expert_loads,
        )
        gate, up = gate_up_output.chunk(2, dim=-1)
        hidden = functional.silu(gate) * up
        return self._dry_run_grouped_matmul(
            hidden, down.transpose(-2, -1), expert_loads,
        )

    def _differentiable_all_to_all(
1879
1880
1881
1882
1883
1884
1885
1886
1887
1888
1889
1890
1891
1892
1893
1894
1895
1896
1897
            self, input_tensor: Any,
            input_splits: tuple[int, ...], output_splits: tuple[int, ...],
    ) -> Any:
        """Execute the real differentiable EP collective with configured splits."""
        from hyper_parallel.platform import get_platform  # pylint: disable=C0415

        try:
            ep_group = self._ep_mesh().get_group()
        except (KeyError, TypeError, AttributeError) as error:
            if len(input_splits) == 1 and input_splits == output_splits:
                return input_tensor.clone()
            raise ValueError(
                "Configured MoE dry-run requires an accessible EP process group"
            ) from error
        return get_platform().differentiable_all_to_all_single(
            input_tensor,
            list(input_splits),
            list(output_splits),
            group=ep_group,
1898
1899
1900
1901
1902
1903
1904
1905
1906
1907
1908
1909
1910
1911
1912
1913
1914
1915
1916
1917
1918
1919
1920
1921
1922
1923
1924
1925
1926
1927
1928
1929
1930
1931
        )

    def _ep_group_identity(self) -> tuple[int, int]:
        """Return EP group size/rank, falling back to the configured topology."""
        try:
            ep_mesh = self._ep_mesh()
            return int(ep_mesh.size()), int(ep_mesh.get_local_rank())
        except (KeyError, TypeError, AttributeError):
            return int(self._base.mesh.ep_size), int(self._base.mesh.ep_rank)

    def _ep_mesh(self) -> Any:
        """Return the expert-parallel child mesh from the HyperModels mesh context."""
        expert_mesh = getattr(self._base.mesh, "fsdp_moe_mesh", None)
        if expert_mesh is None:
            raise ValueError("Configured MoE dry-run requires an expert mesh")
        return expert_mesh["ep"]

    def _restore_moe(self) -> None:
        """Restore every MoE instance forward replaced by this scope."""
        if self._ep_compute_module is not None:
            self._ep_compute_module.ep_routed_forward = self._original_ep_compute
            self._ep_compute_module = None
            self._original_ep_compute = None
        self._moe_module_names.clear()
        while self._patched_moe_modules:
            module, had_forward, original_forward = self._patched_moe_modules.pop()
            if had_forward:
                module.forward = original_forward
            else:
                del module.forward


def _report_limitations(metadata: Optional[Dict[str, Any]] = None) -> list[str]:
    """Return static limitations plus run-specific simulation caveats."""
1945
1946
1947
1948
1949
1950
1951
1952
1953
1954
def _category_name(category: Any) -> str:
    """Return a stable JSON key for a MemTracker category."""
    if isinstance(category, str):
        return category
    value = getattr(category, "value", None)
    return str(value if value is not None else category)


def _normalize_snapshot(snapshot: Dict[Any, Dict[Any, int]]) -> Dict[str, Dict[str, int]]:
    """Convert device/category objects in a MemTracker snapshot to JSON keys."""
1974
1975
1976
1977
1978
1979
1980
1981
1982
    """Aggregate category bytes for devices matching ``device_type``."""
    breakdown: Dict[str, int] = {}
    for device, categories in snapshot.items():
        if device.split(":", maxsplit=1)[0] != device_type:
            continue
        for category, value in categories.items():
            if category == "Total":
                continue
            breakdown[category] = breakdown.get(category, 0) + int(value)
2078
2079
2080
2081
2082
2083
2084
2085
2086
2087

        @staticmethod
        def _is_fake_dtensor_nonleaf(parameter: Any) -> bool:
            """Return whether a failed hook belongs to a fake DTensor view."""
            local_tensor = getattr(parameter, "_local_tensor", None)
            return bool(
                getattr(parameter, "_is_fake_wrapper", False)
                and local_tensor is not None
                and not bool(getattr(local_tensor, "is_leaf", True))
            )
2103
2104
2105
2106
2107
2108
2109
2110
2111
                winfos = self._update_and_maybe_create_winfos(parameter, _MemRefType.PARAM)
                parameter_memory += sum(winfo.mem_consumed for winfo in winfos)
                parameter_gradient = next(self._gradient_candidates(parameter), None)
                if parameter_gradient is not None:
                    self._update_and_maybe_create_winfos(parameter_gradient, _MemRefType.GRAD)
                if self._param_to_grad_hook_handles.get(parameter) is not None or not install_grad_hooks:
                    continue
                grad_hook_handle = parameter.register_hook(_grad_hook)
                try:
2111
2112
2113
2114
2115
2116
2117
2118
2119
2120
2121
2122
                try:
                    post_accumulate_hook_handle = parameter.register_post_accumulate_grad_hook(
                        lambda param: _grad_hook(param.grad),
                    )
                except RuntimeError as error:
                    if "non-leaf" not in str(error) or not self._is_fake_dtensor_nonleaf(parameter):
                        grad_hook_handle.remove()
                        raise
                    post_accumulate_hook_handle = self._NoOpHookHandle()
                self._param_to_grad_hook_handles[parameter] = (
                    grad_hook_handle, post_accumulate_hook_handle,
                )
2122
2123
2124
2125
2126
2127
2128
2129
2130
2131
                )

            buffer_memory = 0
            for buffer in module.buffers():
                winfos = self._update_and_maybe_create_winfos(buffer, _MemRefType.BUFFER)
                buffer_memory += sum(winfo.mem_consumed for winfo in winfos)
            return parameter_memory, buffer_memory

        def __init__(self) -> None:
            """Initialize base tracking state and empty module indexes."""
2162
2163
2164
2165
2166
2167
2168
2169
2170
                if local_tensor is not None and local_tensor.is_leaf
                else None
            )
            if local_gradient is not None and id(local_gradient) not in seen:
                yield local_gradient

        @staticmethod
        def _module_gradient_parameters(module: Any) -> Iterator[Any]:
            """Yield module and FSDP-managed parameters without duplicate identities."""
2176
2177
2178
2179
2180
2181
2182
2183
2184
2185
2186
2187
2188
            for submodule in module.modules():
                scheduler = getattr(submodule, "hsdp_scheduler", None)
                state = getattr(scheduler, "hsdp_state", None)
                for hsdp_param in getattr(state, "hsdp_params", ()):
                    for attribute in ("sharded_param", "unsharded_param"):
                        parameter = getattr(hsdp_param, attribute, None)
                        if parameter is not None and id(parameter) not in seen:
                            seen.add(id(parameter))
                            yield parameter

        def refresh_parameter_gradients(self, module: Any) -> list[Any]:
            """Classify all currently retained local parameter gradients."""
            import torch  # pylint: disable=C0415
2208
2209
2210
2211
2212
2213
2214
2215
2216
                    self._fake_dtensor_grad_bridge_refs[id(parameter)] = weakref.ref(parameter)
            for parameter in self._module_gradient_parameters(module):
                for gradient in self._gradient_candidates(parameter):
                    if id(gradient) in seen:
                        continue
                    seen.add(id(gradient))
                    self._update_and_maybe_create_winfos(gradient, _MemRefType.GRAD)
                    gradients.append(gradient)
            return gradients
2220
2221
2222
2223
2224
2225
2226
2227
2228
2229
2230
2231
2232
2233
2234
2235
2236
2237
2238
2239
2240
2241
            stale_parameter_ids = []
            for parameter_id, parameter_ref in self._fake_dtensor_grad_bridge_refs.items():
                parameter = parameter_ref()
                if parameter is None:
                    stale_parameter_ids.append(parameter_id)
                    continue
                parameter._dry_run_grad_bridge = None
            for parameter_id in stale_parameter_ids:
                self._fake_dtensor_grad_bridge_refs.pop(parameter_id, None)

        def _remove_from_fqn(self, module_fqn: str, module: Any) -> None:
            """Remove a module from one FQN bucket and prune an empty bucket."""
            bucket = self._module_stats_by_fqn.get(module_fqn)
            if bucket is None:
                return
            bucket.pop(module, None)
            if not bucket:
                self._module_stats_by_fqn.pop(module_fqn, None)

        def _rebind_module_stats(self, module: Any) -> None:
            """Bind a module's latest stats FQN while preserving first-seen order."""
            module_stats = self.memory_tracking.get(module)
2239
2240
2241
2242
2243
2244
2245
2246
2247
2248
2249
2250
2251
2252
2253
        def _rebind_module_stats(self, module: Any) -> None:
            """Bind a module's latest stats FQN while preserving first-seen order."""
            module_stats = self.memory_tracking.get(module)
            if module_stats is None:
                return
            new_fqn = str(module_stats.mod_fqn)
            old_fqn = self._module_fqn_by_module.get(module)
            if old_fqn == new_fqn:
                return
            if old_fqn is not None:
                self._remove_from_fqn(old_fqn, module)
            if module not in self._module_registration_order:
                self._module_registration_order[module] = self._next_module_registration_order
                self._next_module_registration_order += 1
            self._module_stats_by_fqn.setdefault(new_fqn, {})[module] = module_stats
2262
2263
2264
2265
2266
2267
2268
2269
2270
            )

            module_fqn = self._mod_tracker.get_known_fqn(module)
            if module_fqn is None:
                raise RuntimeError("MemTracker could not resolve a module FQN")
            previous_peaks = {}
            if module not in self.memory_tracking:
                module_stats = _ModMemStats(module_fqn)
                parameter_mem, buffer_mem = self._track_module_params_and_buffers(
2275
2276
2277
2278
2279
2280
2281
2282
2283
2284
2285
2286
2287
2288
2289
2290
2291
2292
2293
2294
                module_stats.buffer_mem = buffer_mem
                module_stats.input_mem = self._track_inputs_or_outputs(inputs)
                self.memory_tracking[module] = module_stats
                state = _ModState.PRE_FW
            elif self._mod_tracker.is_bw:
                module_stats = self.memory_tracking[module]
                state = _ModState.PRE_FW_AC
                if getattr(self, "_ac_mod", None) is None:
                    self._ac_mod = weakref.ref(module)
                    self._in_ac = True
            else:
                module_stats = self.memory_tracking[module]
                state = _ModState.PRE_FW
                previous_peaks = dict(module_stats.local_peak)
                module_stats.mod_fqn = module_fqn
                module_stats.input_mem = self._track_inputs_or_outputs(inputs)

            memory_snapshot = self.get_tracker_snapshot()
            if state == _ModState.PRE_FW:
                module_stats.local_peak = {
2307
2308
2309
2310
2311
2312
2313
2314
2315
2316
2317
2318
2319
2320
            self._rebind_module_stats(module)

        def reset_mod_stats(self) -> None:
            """Clear module statistics and all indexes derived from them."""
            super().reset_mod_stats()
            self._module_stats_by_fqn.clear()
            self._module_fqn_by_module.clear()
            self._module_registration_order.clear()
            self._next_module_registration_order = 0
            self._fake_dtensor_grad_bridge_refs.clear()

        def _update_peak_stats(self, peak_state: Any) -> None:
            """Update active module peaks in registration order and global peak."""
            active_modules = {}
2327
2328
2329
2330
2331
2332
2333
2334
2335
            )
            current_snapshot = self._curr_mem_snap
            for _, module_stats in ordered_modules:
                if peak_state not in module_stats.snapshots:
                    continue
                for device, device_snapshot in current_snapshot.items():
                    if module_stats.local_peak.get(device, 0) < device_snapshot["Total"]:
                        module_stats.local_peak[device] = device_snapshot["Total"]
                        module_stats.snapshots[peak_state][-1][device] = copy.deepcopy(device_snapshot)
2358
2359
2360
2361
2362
2363
2364
2365
2366
2367
2368
2369
2370
2371

def _operator_phase(tracker: Any) -> str:
    """Resolve the current training phase from MemTracker state."""
    if getattr(tracker, "_in_opt", False):
        return "optimizer"
    from hyper_parallel.core.activation_checkpoint.recompute_state import (  # pylint: disable=C0415
        is_recomputing,
    )
    if is_recomputing():
        return "recompute"
    module_tracker = getattr(tracker, "_mod_tracker", None)
    if getattr(module_tracker, "is_bw", False):
        return "backward"
    return "forward"
2394
2395
2396
2397
2398
2399
2400
2401
2402
2403
            if not isinstance(value, torch.Tensor):
                continue
            try:
                tensor_storages = get_untyped_storages(value)
            except (RuntimeError, TypeError):
                continue
            for storage in tensor_storages:
                storage_key = int(getattr(storage, "_cdata", id(storage)))
                storages[storage_key] = {
                    "device": str(value.device),
2409
2410
2411
2412
2413
2414
2415
2416
2417
    def _lookup_tracked_storage(storage: Any) -> Optional[Dict[str, Any]]:
        """Read MemTracker metadata for one storage without scanning all entries."""
        entry = tracker._WINFO.get(storage)
        if entry is None:
            return None
        winfo, storage_ref = entry
        return {
            "device": str(winfo.device),
            "size": int(winfo.mem_consumed),
2455
2456
2457
2458
2459
2460
2461
2462
2463
2464
2465
            super().__enter__()
            try:
                self._install_storage_resize_hook()
                self._seed_existing_state()
            except Exception:
                super().__exit__(None, None, None)
                raise
            return self

        def __exit__(self, *args: Any) -> Any:
            """Restore MemTracker's storage resize hook before leaving dispatch mode."""
2473
2474
2475
2476
2477
2478
2479
2480
2481
2482
2483
2484
2485
2486

            @functools.wraps(previous_resize)
            def resize_(storage: Any, size: int) -> Any:
                """Resize storage through MemTracker, then record its new lifetime."""
                old_size = int(storage.size())
                result = previous_resize(storage, size)
                new_size = int(storage.size())
                if old_size != new_size:
                    self._record_storage_resize(storage, old_size, new_size)
                return result

            torch.UntypedStorage.resize_ = resize_  # type: ignore[method-assign, assignment]

        def _restore_storage_resize_hook(self) -> None:
2510
2511
2512
2513
2514
2515
2516
2517
2518
                    "storage": storage,
                }
                tracked_storages[storage_key] = tracked_storage
            if not output_storages:
                return
            task_index = self._next_task_index
            self._next_task_index += 1
            devices = {storage["device"] for storage in output_storages.values()}
            self._track_new_outputs(
2536
2537
2538
2539
2540
2541
2542
2543
            """Apply MemTracker role changes to active CSV storage blocks."""
            for storage_key, storage_info in _storage_map(tensors).items():
                block = self._active_blocks.get(storage_key)
                if block is None:
                    continue
                tracked_storage = _lookup_tracked_storage(storage_info["storage"])
                if tracked_storage is not None:
                    block["type"] = tracked_storage["type"]
2546
2547
2548
2549
2550
2551
2552
2553
2554
            """Close active blocks and discard their lifetime bookkeeping."""
            for storage_key in sorted(storage_keys):
                block = self._active_blocks.pop(storage_key, None)
                if block is None:
                    continue
                self._active_storage_refs.pop(storage_key, None)
                self._active_storage_sizes.pop(storage_key, None)
                self._released_storage_keys.discard(storage_key)
                block["end_time_stamp"] = end_index
2617
2618
2619
2620
2621
2622
2623
2624
2625
2626
2627
2628
                    if storage is not None
                    else None
                )
                if tracked_storage is None or tracked_storage["size"] <= 0:
                    block["end_time_stamp"] = self._next_task_index
                    block["_end_lifecycle_event_index"] = self._next_lifecycle_event_index
                    self._next_lifecycle_event_index += 1
                    continue
                block["type"] = tracked_storage["type"]
                block["end_time_stamp"] = _PERSISTENT_END_INDEX
                block["is_persistent"] = 1
            self._active_blocks.clear()
2712
2713
2714
2715
2716
2717
2718
2719
2720
                self._active_storage_sizes[storage_key] = output_storage["size"]
                storage_ref = tracked_storage.get("storage_ref")
                storage = storage_ref() if storage_ref is not None else None
                if storage is None:
                    storage = output_storage["storage"]
                self._active_storage_refs[storage_key] = weakref.ref(
                    storage,
                    lambda _, key=storage_key: self._released_storage_keys.add(key),
                )
2720
2721
2722
2723
2724
2725
2726
2727
2728
2729
2730
2731
2732
2733
2734
2735
2736
2737
2738
2739
2740
2741
2742
2743
                )

        def _record_storage_resize(self, storage: Any, old_size: int, new_size: int) -> None:
            """Record one tracked storage capacity transition as a block lifetime."""
            storage_key = int(getattr(storage, "_cdata", id(storage)))
            tracked_storage = _lookup_tracked_storage(storage)
            if tracked_storage is None:
                return
            device = tracked_storage["device"]
            if device.split(":", maxsplit=1)[0] != device_type:
                return

            if old_size > 0 or storage_key in self._active_blocks:
                self._close_blocks({storage_key}, self._next_task_index)
            if new_size <= 0:
                return

            task_index = self._next_task_index
            self._next_task_index += 1
            output_storages = {
                storage_key: {
                    "device": device,
                    "size": new_size,
                    "storage": storage,
2742
2743
2744
2745
2746
2747
2748
2749
2750
2751
                    "size": new_size,
                    "storage": storage,
                },
            }
            tracked_storages = {storage_key: tracked_storage}
            self._track_new_outputs(
                {storage_key},
                output_storages,
                tracked_storages,
                {
2815
2816
2817
2818
2819
2820
2821
2822
2823
2824
2825
2826

def _report_device_name(device: str, device_type: str, report_device_type: Optional[str]) -> str:
    """Replace a simulation-device prefix with its logical report prefix."""
    if report_device_type is None:
        return device
    prefix, separator, suffix = device.partition(":")
    if prefix != device_type:
        return device
    return report_device_type + (separator + suffix if separator else "")


def _rewrite_module_snapshot_devices(
2829
2830
2831
2832
2833
2834
2835
2836
2837
        report_device_type: Optional[str],
) -> None:
    """Rewrite device keys in every serialized module snapshot in place."""
    if report_device_type is None or report_device_type == device_type:
        return
    for module in modules:
        for snapshots in module["snapshots"].values():
            for snapshot in snapshots:
                renamed = {
2869
2870
2871
2872
2873
2874
2875
2876
    peak = _normalize_snapshot(tracker.get_tracker_snapshot("peak"))
    current = _normalize_snapshot(tracker.get_tracker_snapshot("current"))
    peak_bytes = _snapshot_total(peak, device_type)
    if peak_bytes <= 0:
        raise ValueError(
            f"MemTracker observed no {device_type!r} tensor memory; "
            f"tracked devices are {sorted(peak)}"
        )
2908
2909
2910
2911
2912
2913
2914
2915
2916
2917
2918
2919
2920
2921
2922
2923
2924
2925
2926
2927
2928


def write_memory_report(report: Dict[str, Any], output_path: str) -> str:
    """Atomically write the complete JSON report including metadata."""
    destination = Path(output_path).expanduser().resolve()
    destination.parent.mkdir(parents=True, exist_ok=True)
    descriptor, temporary_path = tempfile.mkstemp(
        prefix=f".{destination.name}.",
        suffix=".tmp",
        dir=str(destination.parent),
    )
    try:
        with os.fdopen(descriptor, "w", encoding="utf-8") as report_file:
            json.dump(report, report_file, indent=2, sort_keys=True)
            report_file.write("\n")
        os.replace(temporary_path, destination)
    except Exception:
        if os.path.exists(temporary_path):
            os.unlink(temporary_path)
        raise
    return str(destination)
hyper_parallel/platform/torch/dtensor.py
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
    @staticmethod
    def _has_function_level_dispatch(func: Any) -> bool:
        """Return whether a FakeTensor wrapper call has a Python-level rule."""
        # pylint: disable=C0415
        from torch._ops import OpOverload, OpOverloadPacket
        from hyper_parallel.core.shard._op_dispatch import _OP_DISPATCHER
        from hyper_parallel.core.shard.ops.parallel_ops_register import get_distributed_op
        from hyper_parallel.core.tensor_parallel._ce_op_registry import is_loss_parallel_op

        if isinstance(func, (OpOverload, OpOverloadPacket)):
            return False
        op_name = getattr(func, "__name__", "")
        return (
            op_name in _OP_DISPATCHER.layout_infer_ops
            or get_distributed_op(op_name) is not None
            or op_name in _OP_DISPATCHER._random_ops  # pylint: disable=W0212
            or is_loss_parallel_op(op_name)
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
            shape: Optional logical global tensor shape.
        """
        if isinstance(local_tensor, DTensorBase):
            # Copy from existing DTensorBase — use alias_placements to preserve multi-axis ordering
            if getattr(local_tensor, "_is_fake_wrapper", False):
                copy_placements = (
                    local_tensor.layout.alias_placements
                    if local_tensor.layout
                    else local_tensor.placements
                )
                return cls(
                    local_tensor._local_tensor,
                    local_tensor.device_mesh,
                    copy_placements,
                    local_tensor.layout,
85
86
87
88
89
90
91
92
93
        # FakeTensor is itself a Python Tensor subclass, so it must be wrapped.
        # pylint: disable=C0415
        from torch._subclasses.fake_tensor import FakeTensor
        if isinstance(local_tensor, FakeTensor):
            t = Tensor._make_wrapper_subclass(
                cls,
                local_tensor.size(),
                strides=local_tensor.stride(),
                storage_offset=local_tensor.storage_offset(),
 95
 96
 97
 98
 99
100
101
102
103
                layout=local_tensor.layout,
                device=torch.device("meta"),
                requires_grad=local_tensor.requires_grad,
            )
            t._is_fake_wrapper = True
        else:
            # Create Tensor subclass instance, sharing local_tensor's underlying storage.
            t = Tensor._make_subclass(cls, local_tensor, local_tensor.requires_grad)
        t.__init_data__(local_tensor, device_mesh, placements, layout, shape)
136
137
138
139
140
141
142
143
            and getattr(arg, "_is_fake_wrapper", False)
            for arg in flat_args
        )
        if has_fake_wrapper and not cls._has_function_level_dispatch(func):
            return super().__torch_function__(func, types, args, kwargs)
        from hyper_parallel.core.shard._op_dispatch import _OP_DISPATCHER
        out = _OP_DISPATCHER.dispatch(func, args, kwargs)
        return out
145
146
147
148
149
150
151
152
153
154
155
156
157
158
    @classmethod
    def __torch_dispatch__(cls, func, types, args=(), kwargs=None):
        """Dispatch FakeTensor-wrapper DTensors through layout rules."""
        # pylint: disable=C0415
        from hyper_parallel.core.shard._op_dispatch import _OP_DISPATCHER
        return _OP_DISPATCHER.dispatch(func, args, kwargs or {})

    def __tensor_flatten__(self):
        """Expose local storage to FakeTensor and MemTracker machinery."""
        context = (
            self._device_mesh,
            self._alias_placements(),
            self._layout,
            getattr(self, "_global_shape", None),
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
            self._alias_placements(),
            self._layout,
            getattr(self, "_global_shape", None),
        )
        return ["_local_tensor"], context

    @classmethod
    def __tensor_unflatten__(cls, inner_tensors, context, outer_size, outer_stride):
        """Rebuild a wrapper DTensor from its flattened local tensor."""
        del outer_size, outer_stride
        device_mesh, placements, layout, shape = context
        return cls(
            inner_tensors["_local_tensor"],
            device_mesh=device_mesh,
            placements=placements,
            layout=layout,
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
        Returns:
            Optional[Tensor]: The gradient tensor, or None if no gradient is set.
        """
        if getattr(self, "_is_fake_wrapper", False):
            outer_grad = Tensor.grad.__get__(self, type(self))  # pylint: disable=C2801
            if outer_grad is not None:
                return outer_grad
            pending_grad = getattr(self, "_fake_pending_grad", None)
            if pending_grad is not None:
                return pending_grad
        if not self._local_tensor.is_leaf and not self._local_tensor.retains_grad:
            return None
        return self._local_tensor.grad

    @grad.setter
    def grad(self, value: Optional[Tensor]) -> None:
199
200
201
202
203
204
205
206
207
208
209
210
        Args:
            value (Optional[Tensor]): The gradient tensor to set, or None to clear.
        """
        if getattr(self, "_is_fake_wrapper", False):
            local_grad = value.to_local() if isinstance(value, DTensorBase) else value
            outer_grad = value
            if value is not None and not isinstance(value, DTensorBase):
                outer_grad = self.__class__(
                    value,
                    device_mesh=self._device_mesh,
                    placements=self._alias_placements(),
                    layout=self._layout,
209
210
211
212
213
214
215
216
217
218
219
220
221
                    placements=self._alias_placements(),
                    layout=self._layout,
                    shape=getattr(self, "_global_shape", None),
                )
            Tensor.grad.__set__(self, outer_grad)  # pylint: disable=C2801
            self._local_tensor.grad = local_grad
            if value is None:
                self._fake_pending_grad = None
            return
        self._local_tensor.grad = value

    @property
    def requires_grad(self) -> bool:
356
357
358
359
360
361
362
363
364
365
366
    @property
    # pylint: disable=C2801
    def data(self):
        """Return the underlying Tensor's data view, bypassing DTensor wrappers."""
        if getattr(self, "_is_fake_wrapper", False):
            return self._local_tensor.data
        return Tensor.data.__get__(self, type(self))

    @data.setter
    # pylint: disable=C2801
    def data(self, value: Tensor) -> None:
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
    def data(self, value: Tensor) -> None:
        """Set the underlying tensor data, extracting the local shard if a DTensor is given."""
        local_value = value.to_local() if isinstance(value, DTensorBase) else value
        if getattr(self, "_is_fake_wrapper", False):
            if self.dtype != local_value.dtype:
                raise ValueError(
                    "Fake-wrapper DTensor data replacement must preserve dtype, "
                    f"got {local_value.dtype} for {self.dtype}"
                )
            self._local_tensor = local_value
            return
        # Tensor.data.__set__ on a Tensor subclass otherwise enters __torch_function__
        # and only rebinds _local_tensor through DTensor dispatch.
        with getattr(torch, "_C").DisableTorchFunctionSubclass():
            Tensor.data.__set__(self, local_value)
483
484
485
486
487
488
489
490
491
492
        Returns:
            DTensorBase: A new DTensor with the converted local tensor.
        """
        new_local = self._local_tensor.to(*args, **kwargs)
        if getattr(self, "_is_fake_wrapper", False):
            return self.__class__(
                new_local,
                device_mesh=self._device_mesh,
                placements=self._alias_placements(),
                layout=self._layout,
hyper_parallel/platform/torch/fully_shard/param.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
79
80


def _register_fake_dtensor_grad_capture(param: DTensor) -> None:
    """Expose an inner FakeTensor leaf gradient through its DTensor wrapper."""
    if not getattr(param, "_is_fake_wrapper", False) or not param._local_tensor.requires_grad:
        return
    local_tensor = param._local_tensor
    capture_tensor_ref = getattr(param, "_fake_grad_capture_tensor_ref", None)
    if capture_tensor_ref is not None and capture_tensor_ref() is local_tensor:
        return
    param_ref = weakref.ref(param)
    local_tensor_ref = weakref.ref(local_tensor)

    def capture_grad(grad: torch.Tensor) -> torch.Tensor:
        captured_param = param_ref()
        if captured_param is not None:
            pending_grad = getattr(captured_param, "_fake_pending_grad", None)
            captured_param._fake_pending_grad = grad if pending_grad is None else pending_grad + grad
        return grad

    def clear_inner_grad(_: torch.Tensor) -> None:
        tensor = local_tensor_ref()
        if tensor is not None:
            tensor.grad = None

    local_tensor.register_hook(capture_grad)
    local_tensor.register_post_accumulate_grad_hook(clear_inner_grad)
    param._fake_grad_capture_tensor_ref = weakref.ref(local_tensor)


def _copy_without_bumping_version(dst: torch.Tensor, src: torch.Tensor) -> None:
    """Copy into ``dst`` while preserving its autograd version counter."""
771
772
773
774
775
776
777
778
779
            unsharded_param,
            requires_grad=self.sharded_param.requires_grad,
        )
        if isinstance(self._unsharded_param, DTensor):
            _register_fake_dtensor_grad_capture(self._unsharded_param)

    def to_sharded(self) -> None:
        self._setattr_on_modules(self.sharded_param)
        if self.unsharded_param_buffers[0] is not self._sharded_param_data:
937
938
939
940
941
942
943
944
945
        if not isinstance(self._sharded_param_data, torch.Tensor):
            return False
        from torch._subclasses.fake_tensor import FakeTensor  # pylint: disable=C0415
        if isinstance(self._sharded_param_data, FakeTensor):
            return self._sharded_param_data.untyped_storage() is local_tensor.untyped_storage()
        sharded_data_ptr = self._sharded_param_data.untyped_storage().data_ptr()
        return (
            # Empty shards may have a zero data pointer and must still be rebuilt.
            sharded_data_ptr > 0
1044
1045
1046
1047
1048
1049
1050
1051
1052
        local_tensor_data = self.sharded_param._local_tensor.data
        local_tensor_storage = local_tensor_data.untyped_storage()
        from torch._subclasses.fake_tensor import FakeTensor  # pylint: disable=C0415
        if isinstance(local_tensor_data, FakeTensor):
            storage_changed = sharded_param_storage is not local_tensor_storage
        else:
            storage_changed = (
                sharded_param_storage.device != local_tensor_storage.device
                or sharded_param_storage.data_ptr() != local_tensor_storage.data_ptr()
hyper_parallel/platform/torch/fully_shard/param_group.py
142
143
144
145
146
147
148
149
150
151
        if self.flat_param_buffer is None:
            return False
        from torch._subclasses.fake_tensor import FakeTensor  # pylint: disable=C0415
        if isinstance(self.flat_param_buffer, FakeTensor):
            flat_storage = self.flat_param_buffer.untyped_storage()
            return all(
                hsdp_param._sharded_param_data.untyped_storage() is flat_storage
                for hsdp_param in self.hsdp_params
            )
        flat_storage_ptr = self.flat_param_buffer.untyped_storage().data_ptr()
hyper_parallel/platform/torch/fully_shard/scheduler.py
117
118
119
120
121
122
123
124
125
126
127
        """Register gradient hooks on all requires-grad outputs to trigger backward pre hook."""
        flat_outputs, _ = tree_flatten(outputs)
        for output in flat_outputs:
            if isinstance(output, torch.Tensor) and output.requires_grad:
                hook_output = output
                if isinstance(output, DTensor) and getattr(output, "_is_fake_wrapper", False):
                    hook_output = output.to_local()
                handle_ref = [None]
                # pylint: disable=C0103, W0102

                def wrapper_for_backward_pre_hook(grad, _handle_ref=handle_ref):
130
131
132
133
134
135
136
137
138
                    if handle is not None:
                        handle.remove()
                    return self._backward_pre_hook(grad)
                # pylint: enable=C0103, W0102
                handle = hook_output.register_hook(wrapper_for_backward_pre_hook)
                handle_ref[0] = handle
        return outputs

    @_dynamo_disable
hyper_parallel/platform/torch/memory_report.py
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86

    Raises:
        ValueError: If the report does not describe a successful run.
    """
    if report.get("status") != "ok":
        raise ValueError("CSV output requires a successful memory report")
    rows = []
    for block in report.get("memory_blocks", []):
        user_tasks = list(block.get("user_tasks", []))
        row = {field: block.get(field, "") for field in _CSV_FIELDS}
        if row["pool_type"] == _FAKE_MEMORY_POOL_TYPE:
            target_device = report.get("metadata", {}).get("target_device", "npu")
            row["pool_type"] = (
                _MEMORY_VISUALIZER_POOL_TYPE
                if target_device == "npu"
                else _CUDA_MEMORY_POOL_TYPE
            )
            user_tasks = []
        row["user_tasks"] = "{" + "-".join(str(task) for task in user_tasks) + "}"
        row["last_user_task"] = user_tasks[-1] if user_tasks else ""
        rows.append(row)
    return rows


def write_memory_csv(report: Dict[str, Any], output_path: str) -> str:
    """Atomically write one memory report as CSV.
 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

    Returns:
        Absolute destination path.
    """
    destination = Path(output_path).expanduser().resolve()
    destination.parent.mkdir(parents=True, exist_ok=True)
    rows = build_memory_csv_rows(report)
    descriptor, temporary_path = tempfile.mkstemp(
        prefix=f".{destination.name}.",
        suffix=".tmp",
        dir=str(destination.parent),
    )
    try:
        with os.fdopen(descriptor, "w", encoding="utf-8", newline="") as report_file:
            writer = csv.DictWriter(report_file, fieldnames=_CSV_FIELDS)
            writer.writeheader()
            writer.writerows(rows)
        os.replace(temporary_path, destination)
    except Exception:
        if os.path.exists(temporary_path):
            os.unlink(temporary_path)
        raise
    return str(destination)


__all__ = ["build_memory_csv_rows", "write_memory_csv"]
hyper_parallel/trainer/config/resolver.py
261
262
263
264
265
266
267
268
269
270
271
272

    config_fields = {field.name: field for field in fields(config_type)}
    unknown = sorted(set(node) - set(config_fields))
    if unknown:
        legacy_dry_run_fields = {"attention", "moe", "tp_cross_entropy"}
        legacy = sorted(set(unknown) & legacy_dry_run_fields)
        if config_type.__name__ == "DryRunConfig" and legacy:
            raise _fail(
                path,
                f"legacy value-dependency fields {legacy} are not supported; "
                "use dry_run.value_dependencies.rules",
            )
hyper_parallel/trainer/dry_run.py
 94
 95
 96
 97
 98
 99
100
101
102
103
104
    """Own all rank-local pipeline chunks for optimizer and memory tracking."""

    def __init__(self, chunks: tuple[_PreparedPipelineChunk, ...]) -> None:
        """Register chunks and expose their shared model configuration."""
        super().__init__()
        self.chunks = nn.ModuleList(chunk.module for chunk in chunks)
        self.config = chunks[0].module.config


class _DryRunBaseTrainer(BaseTrainer):
    """Expose the minimal BaseTrainer initialization used by Dry-run."""
108
109
110
111
112
113
114
115
116
117
118
119
120
121

        Args:
            loss_fn: Optional pipeline-owned loss module.
        """
        if loss_fn is None:
            self._build_loss()
        else:
            self.loss_fn = loss_fn
        self._build_optimizer()
        self._build_training_context()


class HyperModelsDryRunRunner:
    """Execute one shape-accurate fake LLM training step and report memory."""
138
139
140
141
142
143
144
145
        """Require the FakeTensor collective tracking available in PyTorch 2.7."""
        match = re.match(r"^(\d+)\.(\d+)", torch.__version__)
        version = tuple(int(part) for part in match.groups()) if match else (0, 0)
        if version < (2, 7):
            raise RuntimeError(
                "HyperModelsDryRunRunner requires PyTorch >= 2.7, found "
                f"{torch.__version__}"
            )
146
147
148
149
150
151
152
153
154

    def _get_runtime(self) -> torch_dry_run.DryRunRuntime:
        """Return the validated torchrun identity."""
        if self._runtime is None:
            self._runtime = torch_dry_run.DryRunRuntime.from_torchrun_env()
        return self._runtime

    def _validate_config(self) -> DryRunConfig:
        """Validate static Dry-run and unsupported-feature boundaries."""
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
        """Validate static Dry-run and unsupported-feature boundaries."""
        self._check_torch_version()
        dry_run = self.config.dry_run
        if dry_run is None:
            raise ValueError("TrainerConfig.dry_run must be configured")
        if not dry_run.enabled:
            raise ValueError("TrainerConfig.dry_run.enabled must be true")
        if not isinstance(dry_run.output_dir, str) or not dry_run.output_dir.strip():
            raise ValueError("dry_run.output_dir must be a non-empty string")
        if self.config.training.micro_batch_size < 1:
            raise ValueError("training.micro_batch_size must be >= 1")
        activation_checkpoint = self.config.activation_checkpoint.mode
        if activation_checkpoint not in ("off", "none"):
            raise NotImplementedError("LLM Dry-run does not support activation checkpointing")
        if getattr(
                self.config.fsdp_config,
                "activation_checkpointing",
                None,
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
                self.config.fsdp_config,
                "activation_checkpointing",
                None,
        ) not in (False, None, "off", "none"):
            raise NotImplementedError("LLM Dry-run does not support FSDP activation checkpointing")
        if self.config.fsdp_config.enable_offload:
            raise NotImplementedError("LLM Dry-run does not support CPU offload")
        if self.config.peft is not None:
            raise NotImplementedError("LLM Dry-run does not support PEFT model mutation")
        if self.config.compile.enabled:
            raise NotImplementedError("LLM Dry-run does not support torch.compile")
        torch_dry_run._DryRunValueProfile(dry_run)  # pylint: disable=protected-access
        self._validate_topology()
        self._validate_pipeline_config(dry_run)
        return dry_run
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
                "Pipeline Dry-run supports TP, CP, FSDP, and HSDP only; got "
                f"{enabled_unsupported}"
            )
        if accelerator.sequence_parallel or accelerator.loss_parallel:
            raise NotImplementedError(
                "Pipeline Dry-run does not support sequence_parallel or loss_parallel"
            )
        if self.config.activation_swap != "none":
            raise NotImplementedError("Pipeline Dry-run does not support activation swap")
        micro_batch_num = accelerator.pp_micro_batch_num
        if not isinstance(micro_batch_num, int) or isinstance(micro_batch_num, bool) or micro_batch_num < 1:
            raise ValueError("accelerator.pp_micro_batch_num must be a positive integer")
        runtime = self._get_runtime()
        dp_size = runtime.world_size // (accelerator.pp_size * accelerator.tp_size * accelerator.cp_size)
        if self.config.training.global_batch_size % dp_size:
            raise ValueError(
                "training.global_batch_size must be divisible by the PP-local DP size, "
                f"got {self.config.training.global_batch_size} and {dp_size}"
            )
        local_batch_size = self.config.training.global_batch_size // dp_size
222
223
224
225
226
227
228
229
230
231
                f"got {local_batch_size} and {micro_batch_num}"
            )
        normalize_pipeline_schedule(accelerator.pp_schedule, accelerator.pp_vpp)
        if accelerator.pp_layer_split is not None:
            stage_num = accelerator.pp_size * accelerator.pp_vpp
            if len(accelerator.pp_layer_split) != stage_num:
                raise ValueError(
                    f"accelerator.pp_layer_split must contain {stage_num} entries, "
                    f"got {len(accelerator.pp_layer_split)}"
                )
247
248
249
250
251
252
253
254
255
            for name, size in topology_sizes.items()
            if not isinstance(size, int) or isinstance(size, bool) or size < 1
        }
        if invalid_sizes:
            raise ValueError(
                "Dry-run parallel sizes must be positive integers, got "
                f"{invalid_sizes}"
            )
        non_dp_size = accelerator.tp_size * accelerator.cp_size * accelerator.pp_size
253
254
255
256
257
258
259
260
261
                f"{invalid_sizes}"
            )
        non_dp_size = accelerator.tp_size * accelerator.cp_size * accelerator.pp_size
        if runtime.world_size % non_dp_size:
            raise ValueError(
                f"WORLD_SIZE {runtime.world_size} is not divisible by TP*CP*PP "
                f"size {non_dp_size}"
            )
        dp_size = runtime.world_size // non_dp_size
260
261
262
263
264
265
266
267
268
            )
        dp_size = runtime.world_size // non_dp_size
        fsdp_domain_size = dp_size * max(1, accelerator.cp_size)
        if fsdp_domain_size % self.config.fsdp_config.dp_shard_size:
            raise ValueError(
                f"DP+CP size {fsdp_domain_size} is not divisible by "
                f"dp_shard_size {self.config.fsdp_config.dp_shard_size}"
            )
        expert_domain_size = dp_size * accelerator.cp_size * accelerator.tp_size
272
273
274
275
276
277
278
279
                f"ep_size {accelerator.ep_size}"
            )
        expert_dp_size = expert_domain_size // accelerator.ep_size
        if expert_dp_size % self.config.fsdp_config.edp_shard_size:
            raise ValueError(
                f"expert DP size {expert_dp_size} is not divisible by "
                f"edp_shard_size {self.config.fsdp_config.edp_shard_size}"
            )
281
282
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

    @staticmethod
    def _prewarm_standard_meshes(mesh_context: MeshContext) -> None:
        """Materialize child meshes used by the non-pipeline training path."""
        meshes = [mesh_context.device_mesh, mesh_context.fsdp_non_moe_mesh]
        for axis in ("dp", "cp", "tp"):
            meshes.append(mesh_context.device_mesh[axis])
        active_names = tuple(
            axis
            for axis, size in (("cp", mesh_context.cp_size), ("tp", mesh_context.tp_size))
            if size > 1
        )
        if active_names:
            active_mesh = mesh_context.device_mesh[active_names]
            meshes.append(active_mesh)
            meshes.extend(active_mesh[axis] for axis in active_names)
        if mesh_context.dp_cp_mesh is not None:
            meshes.append(mesh_context.dp_cp_mesh)
        dense_selector: str | tuple[str, str] = "fsdp_shard"
        if mesh_context.dp_replicate_size > 1:
            dense_selector = ("fsdp_replicate", "fsdp_shard")
        meshes.append(mesh_context.fsdp_non_moe_mesh[dense_selector])
        if mesh_context.fsdp_moe_mesh is not None:
            meshes.extend((mesh_context.fsdp_moe_mesh, mesh_context.fsdp_moe_mesh["ep"]))
            expert_selector: str | tuple[str, str] = "edp_shard"
            if "edp_replicate" in mesh_context.fsdp_moe_mesh.mesh_dim_names:
                expert_selector = ("edp_replicate", "edp_shard")
            meshes.append(mesh_context.fsdp_moe_mesh[expert_selector])
        for mesh in meshes:
            for mesh_dim in range(mesh.ndim):
                mesh.get_group(mesh_dim)

    def _build_distributed_setup(
            self,
            runtime: torch_dry_run.DryRunRuntime,
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
            runtime: torch_dry_run.DryRunRuntime,
            device_type: str,
    ) -> DistributedSetup:
        """Build the standard non-pipeline mesh and sharding setup."""
        accelerator = self.config.accelerator
        tp_size = accelerator.tp_size
        cp_size = accelerator.cp_size
        dp_size = runtime.world_size // (tp_size * cp_size)
        dp_shard_size = self.config.fsdp_config.dp_shard_size
        dp_replicate_size = dp_size * cp_size // dp_shard_size
        mesh_context = MeshContext(
            dp_size=dp_size,
            dp_replicate_size=dp_replicate_size,
            dp_shard_size=dp_shard_size,
            edp_shard_size=self.config.fsdp_config.edp_shard_size,
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
            ep_size=accelerator.ep_size,
            sequence_parallel=bool(accelerator.sequence_parallel),
            loss_parallel=bool(accelerator.loss_parallel),
        )
        mesh_context.build_meshs(device_type, runtime.world_size)
        mesh_context.dp_rank = mesh_context.device_mesh.get_local_rank("dp")
        mesh_context.tp_rank = mesh_context.device_mesh.get_local_rank("tp")
        mesh_context.cp_rank = mesh_context.device_mesh.get_local_rank("cp")
        mesh_context.ep_rank = (
            mesh_context.fsdp_moe_mesh.get_local_rank("ep")
            if mesh_context.fsdp_moe_mesh is not None
            else 0
        )
        mesh_context.pp_rank = 0
        self._prewarm_standard_meshes(mesh_context)
        fsdp_enabled = (
            dp_shard_size > 1
            or dp_replicate_size > 1
            or self.config.fsdp_config.edp_shard_size > 1
        )
        return DistributedSetup(
            mesh_context=mesh_context,
            strategy_config=self.config.fsdp_config if fsdp_enabled else None,
            plan_overrides=self.config.plan_overrides,
            low_precision_config=getattr(self.config.training, "low_precision", None),
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
        )

    def _select_simulation_device(self, target_device: str) -> str:
        """Select a registered FakeTensor device without probing real hardware."""
        backend_registered = (
            torch.version.cuda is not None
            if target_device == "cuda"
            else hasattr(torch, "npu")
        )
        if not backend_registered:
            logger.info(
                "PyTorch has no %s device guard; using CPU FakeTensor simulation",
                target_device,
            )
            return "cpu"
        try:
            target = torch.device(target_device)
            with FakeTensorMode():
                torch.empty((), device=target)
            return target_device
        except (AttributeError, RuntimeError):
            logger.info(
                "FakeTensor device %s is not registered; using CPU simulation",
                target_device,
            )
        return "cpu"

    def _resolve_target_device(self) -> str:
        """Read the normal-training accelerator type without binding a device."""
        target_device = get_device_type()
        if target_device not in ("cuda", "npu"):
            raise RuntimeError(
                "HyperModels Dry-run requires an available CUDA or NPU runtime; "
                f"detected {target_device!r}"
            )
        self._target_device = target_device
        return target_device

    @staticmethod
    def _init_cpu_data_process_group(runtime: torch_dry_run.DryRunRuntime) -> None:
        """Initialize a temporary Gloo group for the normal data pipeline."""
        if dist.is_initialized():
            raise RuntimeError("Dry-run data probing must start before process-group initialization")
        os.environ.setdefault("GLOO_SOCKET_IFNAME", "lo")
        kwargs = {
            "backend": "gloo",
            "world_size": runtime.world_size,
            "rank": runtime.rank,
        }
        if runtime.world_size == 1:
            kwargs["store"] = dist.HashStore()
        dist.init_process_group(**kwargs)

    @staticmethod
    def _init_fake_process_group(runtime: torch_dry_run.DryRunRuntime) -> None:
        """Initialize PyTorch's fake backend using the torchrun identity."""
        if dist.is_initialized():
            raise RuntimeError("Dry-run must start before any process group is initialized")
        # Import registers the version-matched fake backend factory.
        from torch.testing._internal.distributed import fake_pg  # pylint: disable=C0415,unused-import

        dist.init_process_group(
            backend=dist.Backend.FAKE,
            store=dist.HashStore(),
            world_size=runtime.world_size,
            rank=runtime.rank,
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
            model: nn.Module,
            device: Optional[torch.device] = None,
    ) -> BaseTrainer:
        """Create the BaseTrainer state consumed by shared text-data setup."""
        runtime = self._get_runtime()
        base = BaseTrainer.__new__(BaseTrainer)
        base.config = self.config
        base.local_rank = runtime.local_rank
        base.global_rank = runtime.rank
        base.world_size = runtime.world_size
        base.device = torch.device("cpu") if device is None else device
        base.distributed_setup = setup
        base.mesh = setup.mesh_context
        base.device_mesh = base.mesh.device_mesh
        base.dp_cp_mesh = base.mesh.dp_cp_mesh
        base.model_config = model.config
        if base.config.training.seed is None:
            base.config.training.seed = base.default_seed
        set_seed(
            base.config.training.seed,
            base.config.training.enable_full_determinism,
        )
        return base

    def _read_training_batch(
            self,
            setup: DistributedSetup,
457
458
459
460
461
462
463
464
465
466
467
468
469
470
            model: nn.Module,
            device: Optional[torch.device] = None,
    ) -> PreparedDryRunBatch:
        """Read one CPU batch through the normal TextTrainer data lifecycle."""
        data_runtime = DryRunDataProbe(self._prepare_data_base(setup, model, device))
        try:
            data_runtime.build()
            return data_runtime.read_first_batch()
        finally:
            data_runtime.close()

    def _mock_training_batch(
            self,
            batch: PreparedDryRunBatch,
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
            model: nn.Module,
            tp_size: int,
    ) -> _DryRunTrainingBatch:
        """Erase CPU values and retain only the normal batch structure."""
        valid_count, owned = torch_dry_run.derive_tp_target_counts(
            batch.loss_inputs,
            int(model.config.vocab_size),
            tp_size,
        )
        mocker = torch_dry_run.DryRunBatchMocker(
            fake_mode,
            self._simulation_torch_device(),
        )
        return _DryRunTrainingBatch(
            model_inputs=mocker.mock(batch.model_inputs),
            loss_inputs=mocker.mock(batch.loss_inputs),
            token_counts=dict(batch.token_counts),
            valid_token_count=valid_count,
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524

    @contextmanager
    def _dry_run_fsdp_device(self) -> Iterator[None]:
        """Temporarily resolve FSDP storage to the simulation device."""
        original = fully_shard_api._get_device_from_mesh  # pylint: disable=protected-access
        original_concatenate = DeviceMesh.concatenate
        simulation_device = self._simulation_torch_device()

        def _resolve_device(mesh: Any) -> torch.device:
            del mesh
            return simulation_device

        def _concatenate_meshes(meshes: Any) -> DeviceMesh:
            # Mesh rank maps are host metadata and must never become FakeTensors.
            with unset_fake_temporarily():
                return original_concatenate(meshes)

        fully_shard_api._get_device_from_mesh = _resolve_device  # pylint: disable=protected-access
        DeviceMesh.concatenate = staticmethod(_concatenate_meshes)
        try:
            yield
        finally:
            fully_shard_api._get_device_from_mesh = original  # pylint: disable=protected-access
            DeviceMesh.concatenate = staticmethod(original_concatenate)

    def _simulation_torch_device(self) -> torch.device:
        """Return the rank-local logical device used by FakeTensor execution."""
        if self._simulation_device == "cpu":
            return torch.device("cpu")
        return torch.device(
            self._simulation_device,
            self._get_runtime().local_rank,
        )
542
543
544
545
546
547
548
549
        return model

    def _record_model_dtype(self, model: nn.Module) -> None:
        """Record the effective floating dtype from the normally built model."""
        self._resolved_dtype = next(
            (tensor.dtype for tensor in model.parameters() if tensor.is_floating_point()),
            torch.get_default_dtype(),
        )
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
            model: nn.Module,
            loss_fn: Optional[nn.Module] = None,
    ) -> _DryRunBaseTrainer:
        """Create only the BaseTrainer state needed by one fake training step."""
        runtime = self._get_runtime()
        base = _DryRunBaseTrainer.__new__(_DryRunBaseTrainer)
        base.config = self.config
        base.local_rank = runtime.local_rank
        base.global_rank = runtime.rank
        base.world_size = runtime.world_size
        base.device = self._simulation_torch_device()
        base.distributed_setup = setup
        base.mesh = setup.mesh_context
        base.device_mesh = base.mesh.device_mesh
        base.dp_cp_mesh = base.mesh.dp_cp_mesh
        base.model = model
        base.model_config = model.config
        base.model_parts = [model]
        base.hsdp_model_parts = [
            module for module in model.modules()
            if hasattr(module, "hsdp_scheduler")
        ]
        base.initialize_dry_run_components(loss_fn)
        return base

    def _build_loss(self) -> nn.Module:
        """Build the loss before it is bound into a pipeline stage adapter."""
        loss_fn = self.config.loss_fn.build() if self.config.loss_fn is not None else ModelOutputLoss()
        if not isinstance(loss_fn, nn.Module):
            raise ValueError("config.loss_fn must build a torch.nn.Module")
        return loss_fn

    def _materialize_fake_model(self, model: nn.Module) -> nn.Module:
        """Materialize meta state as FakeTensors without assigning values."""
        fake_parameter = next((
            parameter
            for parameter in model.parameters()
            if hasattr(parameter, "fake_mode")
        ), None)
        if fake_parameter is None:
            # FSDP may move every stage parameter into its internal state, so
            # the plain pipeline root no longer exposes a parameter from which
            # to recover the active FakeTensor converter.
            model.to_empty(device=self._simulation_torch_device())
            return model
        converter = fake_parameter.fake_mode.fake_tensor_converter.meta_converter
        tensor_memo = converter.tensor_memo
        tensor_memo.clear()
        converter.tensor_memo = {}
        try:
            model.to_empty(device=self._simulation_torch_device())
        finally:
            converter.tensor_memo = tensor_memo
        return model

    @staticmethod
    def _initialize_flat_buffers(model: nn.Module) -> None:
        """Materialize enabled zero-copy FSDP flat shards before tracking."""
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
            scheduler = getattr(module, "hsdp_scheduler", None)
            state = getattr(scheduler, "hsdp_state", None)
            if state is None or id(state) in visited_states:
                continue
            visited_states.add(id(state))
            # Match the real forward-pre-hook order: dtype attributes must be
            # initialized before communication buckets choose their dtype.
            state.lazy_init()
            param_group = getattr(state, "param_group", None)
            if param_group is not None and param_group.enable_zero_copy:
                param_group._init_all_gather_buckets()  # pylint: disable=protected-access
                for bucket in param_group.all_gather_buckets:
                    bucket.init_flat_param_buffer(param_group.device)

    @staticmethod
    def _configure_gradient_sync(base: BaseTrainer) -> None:
        """Mark the single simulated micro-step as the final backward."""
628
629
630
631
632
633
634
635
636
637
638
639
    @staticmethod
    def _configure_gradient_sync(base: BaseTrainer) -> None:
        """Mark the single simulated micro-step as the final backward."""
        for model_part in base.hsdp_model_parts:
            model_part.set_requires_gradient_sync(True)
            model_part.set_is_last_backward(True)
            if base.mesh.dp_replicate_size > 1:
                model_part.set_requires_all_reduce(True)

    @staticmethod
    def _loss_context(base: BaseTrainer) -> Any:
        """Return the real loss-parallel context when configured."""
638
639
640
641
642
643
644
645
646
    def _loss_context(base: BaseTrainer) -> Any:
        """Return the real loss-parallel context when configured."""
        if not base.mesh.loss_parallel:
            return nullcontext()
        return loss_parallel(mesh=base.mesh.device_mesh["tp"])

    def _execute_micro_step(
            self,
            base: BaseTrainer,
660
661
662
663
664
665
666
667
668
                labels=batch.loss_inputs.get("labels"),
            )
        del outputs
        if isinstance(loss_value, dict):
            loss_value = torch.stack(list(loss_value.values())).sum()
        with base.model_bwd_context, value_dependencies.logical_scope("backward"):
            loss_value.backward()

        gradient_tensors = tracker.refresh_parameter_gradients(base.model)
811
812
813
814
815
816
817
818
819
                for item in value.values()
                for tensor in cls._tensor_leaves(item)
            )
        if isinstance(value, (list, tuple)):
            return tuple(
                tensor
                for item in value
                for tensor in cls._tensor_leaves(item)
            )
817
818
819
820
821
822
823
824
825
                for item in value
                for tensor in cls._tensor_leaves(item)
            )
        if is_dataclass(value) and not isinstance(value, type):
            return tuple(
                tensor
                for field in fields(value)
                for tensor in cls._tensor_leaves(getattr(value, field.name))
            )
836
837
838
839
840
841
842
843
844
845
846
            }
        if isinstance(value, Mapping):
            return {name: cls._batch_shape_metadata(item) for name, item in value.items()}
        if isinstance(value, (list, tuple)):
            return [cls._batch_shape_metadata(item) for item in value]
        if is_dataclass(value) and not isinstance(value, type):
            return {
                field.name: cls._batch_shape_metadata(getattr(value, field.name))
                for field in fields(value)
            }
        return value
853
854
855
856
857
858
859
860
861
862
863
864
865
866
867
868
869
870
871
872
873
            loss_fn: nn.Module,
            build_state: _DryRunPipelineBuildState,
    ) -> Any:
        """Create the scoped normal-builder adapter used by pipeline Dry-run."""
        num_micro_batches = self.config.accelerator.pp_micro_batch_num

        def build_pipeline_model(
                model: nn.Module,
                request: DeferredModelBuildRequest,
        ) -> nn.Module:
            """Split one raw meta model and apply stage-local infrastructure."""
            if request.distributed_setup is not parallel_context.setup:
                raise RuntimeError("Pipeline Dry-run adapter received an unexpected DistributedSetup")
            if request.sharding_planner is None:
                raise RuntimeError("Pipeline Dry-run requires a normal sharding planner")
            self._record_model_dtype(model)
            chunks = build_pipeline_chunks(
                self.config,
                model,
                loss_fn,
                num_micro_batches,
874
875
876
877
878
879
880
881
882
                parallel_context,
                self._resolved_dtype,
                tuple(int(size) for size in batch.model_inputs["input_ids"].shape),
            )
            chunks = prepare_pipeline_chunks(
                chunks,
                parallel_context.setup,
                request.sharding_planner,
                request.fsdp2_manager,
882
883
884
885
886
887
888
889
890
891
892
893
894
895
896
                request.fsdp2_manager,
                request.validate_placement,
                self._dry_run_fsdp_device,
            )
            pipeline_model = _DryRunPipelineModel(chunks)
            apply_model_init_dtype(pipeline_model, request.model_init_dtype)
            profile.project_model(pipeline_model)
            build_state.chunks = chunks
            return pipeline_model

        return build_pipeline_model

    def _build_pipeline_execution(
            self,
            base: BaseTrainer,
923
924
925
926
927
928
929
930
931
        )
        if stages[-1].is_last_stage:
            labels = batch.loss_inputs.get("labels")
            if labels is None:
                raise ValueError("Pipeline Dry-run requires loss_inputs.labels")
            stages[-1].set_micro_labels(list(labels.chunk(num_micro_batches, dim=0)))
        return stages, schedule, stage_indices, stage_num

    def _pipeline_metadata(
1022
1023
1024
1025
1026
1027
1028
1029
1030
1031
1032
1033
1034
1035
1036
1037
1038
1039
1040
1041
1042
1043
1044
1045
1046
1047
1048
1049
1050
1051
1052
1053
1054
            batch: _DryRunTrainingBatch,
            num_micro_batches: int,
    ) -> _DryRunTrainingBatch:
        """Repeat one normal micro-batch into the configured pipeline batch."""
        micro_batch_size = int(batch.model_inputs["input_ids"].shape[0])

        def repeat(value: Any) -> Any:
            """Repeat tensor leaves carrying the leading batch dimension."""
            if isinstance(value, torch.Tensor) and value.ndim and value.shape[0] == micro_batch_size:
                return torch.cat([value] * num_micro_batches, dim=0)
            if isinstance(value, Mapping):
                return {name: repeat(item) for name, item in value.items()}
            if isinstance(value, list):
                return [repeat(item) for item in value]
            if isinstance(value, tuple):
                return tuple(repeat(item) for item in value)
            if is_dataclass(value) and not isinstance(value, type):
                repeated = replace(value, **{
                    field.name: repeat(getattr(value, field.name))
                    for field in fields(value)
                    if field.init
                })
                for field in fields(value):
                    if not field.init:
                        object.__setattr__(repeated, field.name, repeat(getattr(value, field.name)))
                return repeated
            return value

        return _DryRunTrainingBatch(
            model_inputs=repeat(batch.model_inputs),
            loss_inputs=repeat(batch.loss_inputs),
            token_counts={name: count * num_micro_batches for name, count in batch.token_counts.items()},
            valid_token_count=batch.valid_token_count,
1060
1061
1062
1063
1064
1065
1066
1067
1068
1069
1070
1071
1072
1073
1074
1075
1076
1077
1078
1079
1080
            runtime: torch_dry_run.DryRunRuntime,
            fake_mode: FakeTensorMode,
    ) -> _DryRunTrainingBatch:
        """Read one CPU batch and erase its tensor values for simulation."""
        self._stage = "data_probe"
        data_group_initialized = False
        cpu_setup = None
        try:
            self._init_cpu_data_process_group(runtime)
            data_group_initialized = True
            cpu_setup = self._build_distributed_setup(runtime, "cpu")
            normalize_distributed_setup_overrides(cpu_setup, self.config)
            probe_setup = copy(cpu_setup)
            probe_setup.strategy_config = None
            probe_model = self._build_target_model(probe_setup)
            cpu_batch = self._read_training_batch(cpu_setup, probe_model)
            batch = self._mock_training_batch(
                cpu_batch,
                fake_mode,
                probe_model,
                cpu_setup.mesh_context.tp_size,
1078
1079
1080
1081
1082
1083
1084
1085
1086
1087
1088
1089
1090
1091
1092
                fake_mode,
                probe_model,
                cpu_setup.mesh_context.tp_size,
            )
            del cpu_batch
            del probe_model
            return batch
        finally:
            cpu_setup = None
            if data_group_initialized:
                destroy_distributed_runtime()

    def _run_simulation(
            self,
            runtime: torch_dry_run.DryRunRuntime,
1096
1097
1098
1099
1100
1101
1102
1103
1104
1105
1106
1107
1108
1109
1110
1111
1112
1113
1114
1115
1116
1117
1118
1119
1120
1121
1122
1123
1124
1125
1126
1127
1128
1129
1130
1131
1132
            json_path: str,
            csv_path: str,
    ) -> dict[str, Any]:
        """Execute one fake distributed step and persist its memory report."""
        initialized_here = False
        try:
            self._stage = "distributed_setup"
            self._init_fake_process_group(runtime)
            initialized_here = True
            parallel_context = (
                build_pipeline_distributed_setup(self.config, runtime, self._simulation_device)
                if self.config.accelerator.pp_size > 1
                else None
            )
            setup = (
                parallel_context.setup
                if parallel_context is not None
                else self._build_distributed_setup(runtime, self._simulation_device)
            )
            normalize_distributed_setup_overrides(setup, self.config)

            self._stage = "model_build"
            profile = torch_dry_run._DryRunValueProfile(dry_run)  # pylint: disable=protected-access
            with fake_mode:
                with self._dry_run_fsdp_device():
                    if parallel_context is None:
                        model = self._build_target_model(setup)
                        pipeline_state = None
                        loss_fn = None
                    else:
                        loss_fn = self._build_loss()
                        pipeline_state = _DryRunPipelineBuildState()
                        model = self._build_target_model(
                            setup,
                            self._build_pipeline_adapter(
                                parallel_context,
                                batch,
1134
1135
1136
1137
1138
1139
1140
1141
1142
1143
1144
1145
1146
1147
1148
1149
1150
1151
1152
1153
1154
1155
1156
1157
1158
1159
1160
1161
1162
1163
1164
1165
1166
1167
1168
1169
                                loss_fn,
                                pipeline_state,
                            ),
                        )
                self._record_model_dtype(model)
                if pipeline_state is None:
                    profile.bind_model(model)
                value_dependencies = torch_dry_run.ValueDependencyManager(
                    profile, model, runtime,
                )
                self._stage = "materialize"
                model = self._materialize_fake_model(model)
                profile.configure_tp_cross_entropy_counts(
                    batch.valid_token_count,
                    batch.target_tokens_per_rank,
                )
                if pipeline_state is None:
                    base = self._prepare_base(setup, model)
                    self._stage = "fake_step"
                    report = self._execute_step(base, batch, profile, value_dependencies)
                else:
                    if pipeline_state.chunks is None or parallel_context is None or loss_fn is None:
                        raise RuntimeError("Pipeline Dry-run adapter did not produce local pipeline chunks")
                    if parallel_context.pp_mesh is None:
                        raise RuntimeError("Pipeline Dry-run requires a PP mesh")
                    base = self._prepare_base(setup, model, loss_fn)
                    pipeline_batch = self._repeat_pipeline_batch(
                        batch,
                        self.config.accelerator.pp_micro_batch_num,
                    )
                    self._stage = "pipeline_fake_step"
                    report = self._execute_pipeline_step(
                        base,
                        pipeline_batch,
                        profile,
                        value_dependencies,
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180
1181
1182
1183
1184
1185
1186
1187
1188
1189
1190
1191
1192
1193
1194
1195
1196
1197
1198
1199
1200
1201
1202
1203
1204
1205
1206
                        pipeline_state.chunks,
                        parallel_context.pp_mesh,
                    )

            self._stage = "report_generation"
            torch_dry_run.write_memory_report(report, json_path)
            torch_dry_run.write_memory_csv(report, csv_path)
            logger.info("HyperModels Dry-run memory report: %s", csv_path)
            return report
        except torch_dry_run.UnconfiguredValueDependencyError:
            raise
        except Exception as error:
            raise RuntimeError(
                f"HyperModels Dry-run failed during {self._stage}: {error}"
            ) from error
        finally:
            if initialized_here:
                destroy_distributed_runtime()

    def run(self) -> dict[str, Any]:
        """Execute one fake step, write the rank-local CSV, and return its report."""
        runtime = self._get_runtime()
        dry_run = self._validate_config()
        csv_path = os.path.join(
            dry_run.output_dir,
            f"rank_{runtime.rank}",
            f"rank_{runtime.rank}_memory.csv",
        )
        json_path = os.path.splitext(csv_path)[0] + ".json"
        target_device = self._resolve_target_device()
        self._simulation_device = self._select_simulation_device(target_device)
        fake_mode = FakeTensorMode(allow_non_fake_inputs=True)
        batch = self._probe_training_batch(runtime, fake_mode)
        return self._run_simulation(runtime, dry_run, batch, fake_mode, json_path, csv_path)


__all__ = ["HyperModelsDryRunRunner"]
hyper_parallel/trainer/dry_run_data.py
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
    def read_first_batch(self) -> PreparedDryRunBatch:
        """Read this rank's first fully prepared training micro-batch."""
        dataloader = getattr(self.base, "train_dataloader", None)
        if dataloader is None:
            raise ValueError("A non-empty train dataloader is required")
        self._iterator = iter(dataloader)
        try:
            model_inputs, loss_inputs = self.base.get_batch(self._iterator)
        except StopIteration as exc:
            raise ValueError("Train dataloader produced no batches") from exc
        if not isinstance(model_inputs, Mapping) or not isinstance(loss_inputs, Mapping):
            raise ValueError("dataloader.get_batch must return model_inputs and loss_inputs mappings")
        counts = {
            name: int(value.item())
            for name, value in count_loss_token(dict(loss_inputs)).items()
        }
71
72
73
74
75
76
77
78
79
80
81
82
        return PreparedDryRunBatch(dict(model_inputs), dict(loss_inputs), counts)

    def close(self) -> None:
        """Stop workers owned by the one-shot iterator, if any."""
        iterator = self._iterator
        self._iterator = None
        shutdown_workers = getattr(iterator, "_shutdown_workers", None)
        if callable(shutdown_workers):
            shutdown_workers()  # pylint: disable=not-callable


__all__ = ["DryRunDataProbe", "PreparedDryRunBatch"]
hyper_parallel/trainer/dry_run_pipeline.py
50
51
52
53
54
55
56
57
58
    Raises:
        ValueError: If the schedule or VPP degree is invalid.
    """
    if not isinstance(pp_vpp, int) or isinstance(pp_vpp, bool) or pp_vpp < 1:
        raise ValueError(f"pp_vpp must be a positive integer, got {pp_vpp!r}")
    normalized = (
        "1f1b"
        if schedule_name is None
        else str(schedule_name).lower().replace("-", "_")
58
59
60
61
62
63
64
65
66
67
68
69
70
        else str(schedule_name).lower().replace("-", "_")
    )
    if pp_vpp > 1:
        if normalized not in ("1f1b", "interleaved_1f1b"):
            raise ValueError(
                f"pp_vpp > 1 requires pp_schedule='1f1b', got {schedule_name!r}"
            )
        return "interleaved_1f1b"
    if normalized not in _PIPELINE_SCHEDULES:
        raise ValueError(
            f"Unknown pp_schedule {schedule_name!r}; choose from {sorted(_PIPELINE_SCHEDULES)}"
        )
    return normalized
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114

    Raises:
        ValueError: If ``micro_batch_num`` is not a positive integer.
    """
    if (
            not isinstance(micro_batch_num, int)
            or isinstance(micro_batch_num, bool)
            or micro_batch_num < 1
    ):
        raise ValueError(
            f"micro_batch_num must be a positive integer, got {micro_batch_num!r}"
        )
    normalized = normalize_pipeline_schedule(schedule_name, pp_vpp)
    schedule_class = (
        ScheduleInterleaved1F1B
        if normalized == "interleaved_1f1b"
        else _PIPELINE_SCHEDULES[normalized]
    )
    return schedule_class(stages, micro_batch_num, **schedule_kwargs)


@dataclass(frozen=True)
class DryRunBoundaryLeaf:
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145

    def __post_init__(self) -> None:
        """Validate the static boundary description."""
        if not self.global_shape or any(size <= 0 for size in self.global_shape):
            raise ValueError(
                "DryRunBoundaryLeaf.global_shape must contain positive dimensions, "
                f"got {self.global_shape}"
            )
        if not isinstance(self.dtype, torch.dtype):
            raise ValueError(
                f"DryRunBoundaryLeaf.dtype must be torch.dtype, got {self.dtype!r}"
            )
        if not self.tensor_name:
            raise ValueError("DryRunBoundaryLeaf.tensor_name must be non-empty")
        if self.wire_kind not in ("tensor", "dtensor"):
            raise ValueError(
                "DryRunBoundaryLeaf.wire_kind must be 'tensor' or 'dtensor', "
                f"got {self.wire_kind!r}"
            )
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190

    def __post_init__(self) -> None:
        """Validate the model-builder result at the framework boundary."""
        if not isinstance(self.module, nn.Module):
            raise ValueError("DryRunPipelineChunk.module must be a torch.nn.Module")
        if self.layer_start < 0 or self.layer_end < self.layer_start:
            raise ValueError(
                "DryRunPipelineChunk layer range must be ordered and non-negative, "
                f"got [{self.layer_start}, {self.layer_end})"
            )
        if self.hidden_size <= 0:
            raise ValueError("DryRunPipelineChunk.hidden_size must be positive")
        if self.stage_index < -1:
            raise ValueError("DryRunPipelineChunk.stage_index must be non-negative or -1")
        for unit in self.fsdp_units:
            if not unit or not all(isinstance(module, nn.Module) for module in unit):
                raise ValueError(
                    "DryRunPipelineChunk.fsdp_units must contain non-empty module tuples"
                )

223
224
225
226
227
228
229
230
231
            output_metadata: Ordered metadata used to validate forward send tensors.
            mesh: Existing one-dimensional PP mesh used for peer-rank mapping.
            group: Optional explicit PP process group.
        """
        super().__init__(
            submodule,
            stage_index=stage_index,
            stage_num=stage_num,
            device=device,
231
232
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
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
            device=device,
            group=group,
            mesh=mesh,
        )
        self._input_metadata = input_metadata
        self._output_metadata = output_metadata
        self._meta_cache = []

    def _communicate_meta(self, global_rank: int, meta_send: Any = None) -> Any:
        """Validate sender metadata or supply receiver metadata without communication."""
        del global_rank
        if meta_send is None:
            return [self._input_metadata]
        expected = self._output_metadata
        if len(meta_send) != len(expected):
            raise ValueError(
                "Dry-run PP boundary output count mismatch: expected "
                f"{len(expected)}, got {len(meta_send)}"
            )
        for leaf_index, (actual, expected_leaf) in enumerate(zip(meta_send, expected)):
            if len(actual) != len(expected_leaf):
                raise ValueError(
                    f"Dry-run PP boundary leaf {leaf_index} metadata kind mismatch: "
                    f"expected {len(expected_leaf)} fields, got {len(actual)}"
                )
            if tuple(actual[0]) != tuple(expected_leaf[0]):
                raise ValueError(
                    f"Dry-run PP boundary leaf {leaf_index} shape mismatch: "
                    f"expected {tuple(expected_leaf[0])}, got {tuple(actual[0])}"
                )
            if actual[1] != expected_leaf[1]:
                raise ValueError(
                    f"Dry-run PP boundary leaf {leaf_index} dtype mismatch: "
                    f"expected {expected_leaf[1]}, got {actual[1]}"
                )
            if len(actual) == 4:
                actual_placements = tuple(actual[2].alias_placements)
                expected_placements = tuple(expected_leaf[2].alias_placements)
                if actual_placements != expected_placements:
                    raise ValueError(
                        f"Dry-run PP boundary leaf {leaf_index} layout mismatch: "
                        f"expected {expected_placements}, got {actual_placements}"
                    )
            if bool(actual[-1]) != bool(expected_leaf[-1]):
                raise ValueError(
                    f"Dry-run PP boundary leaf {leaf_index} requires_grad mismatch: "
                    f"expected {bool(expected_leaf[-1])}, got {bool(actual[-1])}"
                )
        return None

    @staticmethod
    def _mock_works(specs: list[tuple[str, Any, int]]) -> list[_NoOpWork]:
        """Return one completed work object for each inherited communication spec."""
        return [_NoOpWork() for _ in specs]

    def exec_fwd_recv_ops(self, micro_index: int) -> list[_NoOpWork]:
        """Allocate and register the inherited forward receive buffers."""
        return self._mock_works(self.fwd_recv_specs(micro_index))

    def exec_fwd_send_ops(self, micro_index: int) -> list[_NoOpWork]:
        """Apply inherited forward-send cache transitions without sending."""
        return self._mock_works(self.fwd_send_specs(micro_index))

    def exec_bwd_recv_ops(self, micro_index: int) -> list[_NoOpWork]:
        """Expose inherited gradient receive buffers without receiving."""
        return self._mock_works(self.bwd_recv_specs(micro_index))

    def exec_bwd_send_ops(self, micro_index: int) -> list[_NoOpWork]:
        """Apply inherited backward-send cache transitions without sending."""
        return self._mock_works(self.bwd_send_specs(micro_index))

    def forward_one_chunk(
            self,
            micro_index: int,
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
            args: Any = None,
            kwargs: Any = None,
    ) -> Any:
        """Select stage-local labels before delegating the forward cache logic."""
        set_micro_index = getattr(self.submodule, "set_micro_index", None)
        if set_micro_index is not None:
            set_micro_index(micro_index)
        return super().forward_one_chunk(micro_index, args, kwargs)

    def set_micro_labels(self, labels: list[torch.Tensor]) -> None:
        """Install last-stage labels split in scheduler micro-batch order."""
        setter = getattr(self.submodule, "set_micro_labels", None)
        if setter is None:
            raise ValueError("The last dry-run PP stage must implement set_micro_labels")
        setter(labels)

    def clear_all_states(self) -> None:
        """Drop all per-run caches and receive-buffer references."""
        self.clear_states()
        self.clear_cache()
        self.last_stage_outputs = None


__all__ = [
    "DryRunBoundaryLeaf",
hyper_parallel/trainer/dry_run_pipeline_assembly.py
84
85
86
87
88
89
90
91
92
        cp_size: int,
        tp_size: int,
) -> tuple[DeviceMesh, DeviceMesh, DeviceMesh]:
    """Build private root, PP, and stage-compute meshes."""
    root_mesh = init_device_mesh(
        device_type=device_type,
        mesh_shape=(pp_size, dp_size, cp_size, tp_size),
        mesh_dim_names=("pp", "dp", "cp", "tp"),
        init_backend=dist.is_initialized(),
 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
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
147
148
149
150
151
152
153
154
        mesh_shape=(pp_size, dp_size, cp_size, tp_size),
        mesh_dim_names=("pp", "dp", "cp", "tp"),
        init_backend=dist.is_initialized(),
    )
    return root_mesh, root_mesh["pp"], root_mesh[("dp", "cp", "tp")]


def _prewarm_pipeline_meshes(root_mesh: DeviceMesh, mesh_context: MeshContext, pp_mesh: DeviceMesh) -> None:
    """Create all process groups before FakeTensorMode or model execution."""
    meshes = [root_mesh, pp_mesh, mesh_context.device_mesh, mesh_context.fsdp_non_moe_mesh]
    for axis in ("dp", "cp", "tp"):
        meshes.append(mesh_context.device_mesh[axis])
    active_names = tuple(
        axis
        for axis, size in (("cp", mesh_context.cp_size), ("tp", mesh_context.tp_size))
        if size > 1
    )
    if active_names:
        active_mesh = mesh_context.device_mesh[active_names]
        meshes.append(active_mesh)
        meshes.extend(active_mesh[axis] for axis in active_names)
    if mesh_context.dp_cp_mesh is not None:
        meshes.append(mesh_context.dp_cp_mesh)
    dense_selector: str | tuple[str, str] = "fsdp_shard"
    if mesh_context.dp_replicate_size > 1:
        dense_selector = ("fsdp_replicate", "fsdp_shard")
    meshes.append(mesh_context.fsdp_non_moe_mesh[dense_selector])
    for mesh in meshes:
        for mesh_dim in range(mesh.ndim):
            mesh.get_group(mesh_dim)


def _prewarm_standard_meshes(mesh_context: MeshContext) -> None:
    """Materialize child meshes used by the existing non-pipeline path."""
    meshes = [mesh_context.device_mesh, mesh_context.fsdp_non_moe_mesh]
    for axis in ("dp", "cp", "tp"):
        meshes.append(mesh_context.device_mesh[axis])
    active_names = tuple(
        axis
        for axis, size in (("cp", mesh_context.cp_size), ("tp", mesh_context.tp_size))
        if size > 1
    )
    if active_names:
        active_mesh = mesh_context.device_mesh[active_names]
        meshes.append(active_mesh)
        meshes.extend(active_mesh[axis] for axis in active_names)
    if mesh_context.dp_cp_mesh is not None:
        meshes.append(mesh_context.dp_cp_mesh)
    dense_selector: str | tuple[str, str] = "fsdp_shard"
    if mesh_context.dp_replicate_size > 1:
        dense_selector = ("fsdp_replicate", "fsdp_shard")
    meshes.append(mesh_context.fsdp_non_moe_mesh[dense_selector])
    if mesh_context.fsdp_moe_mesh is not None:
        meshes.extend((mesh_context.fsdp_moe_mesh, mesh_context.fsdp_moe_mesh["ep"]))
        expert_selector: str | tuple[str, str] = "edp_shard"
        if "edp_replicate" in mesh_context.fsdp_moe_mesh.mesh_dim_names:
            expert_selector = ("edp_replicate", "edp_shard")
        meshes.append(mesh_context.fsdp_moe_mesh[expert_selector])
    for mesh in meshes:
        for mesh_dim in range(mesh.ndim):
            mesh.get_group(mesh_dim)


def _build_standard_distributed_setup(
        config: TrainerConfig,
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
        runtime: DryRunRuntime,
        device_type: str,
) -> _DryRunParallelContext:
    """Build the established non-pipeline MeshContext domains."""
    accelerator = config.accelerator
    tp_size = max(1, accelerator.tp_size)
    cp_size = max(1, accelerator.cp_size)
    ep_size = max(1, accelerator.ep_size)
    dp_size = runtime.world_size // (tp_size * cp_size)
    dp_shard_size = config.fsdp_config.dp_shard_size
    dp_replicate_size = dp_size * cp_size // dp_shard_size
    mesh_context = MeshContext(
        dp_size=dp_size,
        dp_replicate_size=dp_replicate_size,
        dp_shard_size=dp_shard_size,
        edp_shard_size=config.fsdp_config.edp_shard_size,
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
        ep_size=ep_size,
        sequence_parallel=bool(accelerator.sequence_parallel),
        loss_parallel=bool(accelerator.loss_parallel),
    )
    mesh_context.build_meshs(device_type, runtime.world_size)
    mesh_context.dp_rank = mesh_context.device_mesh.get_local_rank("dp")
    mesh_context.tp_rank = mesh_context.device_mesh.get_local_rank("tp")
    mesh_context.cp_rank = mesh_context.device_mesh.get_local_rank("cp")
    mesh_context.ep_rank = (
        mesh_context.fsdp_moe_mesh.get_local_rank("ep")
        if mesh_context.fsdp_moe_mesh is not None
        else 0
    )
    mesh_context.pp_rank = 0
    _prewarm_standard_meshes(mesh_context)
    fsdp_enabled = (
        dp_shard_size > 1
        or dp_replicate_size > 1
        or config.fsdp_config.edp_shard_size > 1
    )
    setup = DistributedSetup(
        mesh_context=mesh_context,
        strategy_config=config.fsdp_config if fsdp_enabled else None,
        plan_overrides=config.plan_overrides,
        low_precision_config=getattr(config.training, "low_precision", None),
197
198
199
200
201
202
203
204
205
        plan_overrides=config.plan_overrides,
        low_precision_config=getattr(config.training, "low_precision", None),
        fp32_main_params=config.optimizer.fp32_main_params,
    )
    return _DryRunParallelContext(setup, None, None)


def build_distributed_setup(
        config: TrainerConfig,
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
        runtime: DryRunRuntime,
        device_type: str,
) -> _DryRunParallelContext:
    """Build root-derived stage meshes without changing public MeshContext semantics."""
    accelerator = config.accelerator
    if max(1, accelerator.pp_size) == 1:
        return _build_standard_distributed_setup(config, runtime, device_type)
    tp_size = max(1, accelerator.tp_size)
    cp_size = max(1, accelerator.cp_size)
    pp_size = max(1, accelerator.pp_size)
    ep_size = max(1, accelerator.ep_size)
    dp_size = runtime.world_size // (tp_size * cp_size * pp_size)
    dp_shard_size = config.fsdp_config.dp_shard_size
    dp_replicate_size = dp_size * cp_size // dp_shard_size
    mesh_context = MeshContext(
        dp_size=dp_size,
        dp_replicate_size=dp_replicate_size,
        dp_shard_size=dp_shard_size,
        edp_shard_size=config.fsdp_config.edp_shard_size,
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
        ep_size=ep_size,
        sequence_parallel=bool(accelerator.sequence_parallel),
        loss_parallel=bool(accelerator.loss_parallel),
    )
    root_mesh, pp_mesh, stage_mesh = _build_pipeline_meshes(
        device_type, pp_size, dp_size, cp_size, tp_size
    )
    mesh_context.build_meshs(device_type, runtime.world_size)
    if tuple(mesh_context.device_mesh.rank_list) != tuple(stage_mesh.rank_list):
        raise ValueError("Pipeline root mesh and MeshContext stage ranks do not match")
    mesh_context.dp_rank = stage_mesh.get_local_rank("dp")
    mesh_context.tp_rank = stage_mesh.get_local_rank("tp")
    mesh_context.cp_rank = stage_mesh.get_local_rank("cp")
    mesh_context.pp_rank = pp_mesh.get_local_rank()
    mesh_context.ep_rank = 0
    _prewarm_pipeline_meshes(root_mesh, mesh_context, pp_mesh)
    fsdp_enabled = dp_shard_size > 1 or dp_replicate_size > 1
    setup = DistributedSetup(
        mesh_context=mesh_context,
        strategy_config=config.fsdp_config if fsdp_enabled else None,
        plan_overrides=config.plan_overrides,
        low_precision_config=getattr(config.training, "low_precision", None),
248
249
250
251
252
253
254
255
256
        plan_overrides=config.plan_overrides,
        low_precision_config=getattr(config.training, "low_precision", None),
        fp32_main_params=config.optimizer.fp32_main_params,
    )
    return _DryRunParallelContext(setup, root_mesh, pp_mesh)


def build_pipeline_chunks(
        config: TrainerConfig,
262
263
264
265
266
267
268
269
270
        micro_batch_shape: tuple[int, int],
) -> tuple[DryRunPipelineChunk, ...]:
    """Invoke and validate the configured model-specific PP builder."""
    if parallel_context.pp_mesh is None:
        raise ValueError("Pipeline chunk construction requires pp_size > 1")
    pp_rank = parallel_context.pp_mesh.get_local_rank()
    # ParallelBatch has already sharded sequence tensors across CP ranks. The
    # boundary contract is global, so restore its global sequence dimension
    # before the planner resolves the stage-local boundary exactly once.
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
        sequence_length=global_sequence_length,
        boundary_dtype=boundary_dtype,
    )
    if not isinstance(chunks, tuple) or not all(isinstance(chunk, DryRunPipelineChunk) for chunk in chunks):
        raise ValueError(
            "dry_run.pipeline_stage_builder must return a tuple of "
            f"DryRunPipelineChunk objects, got {type(chunks).__name__}"
        )
    if len(chunks) != config.accelerator.pp_vpp:
        raise ValueError(
            "Dry-run pipeline builder must return one chunk per virtual stage, "
            f"got {len(chunks)} for pp_vpp={config.accelerator.pp_vpp}"
        )
    if len({chunk.hidden_size for chunk in chunks}) != 1:
        raise ValueError("All dry-run pipeline chunks must use the same hidden_size")
    expected_indices = tuple(
        pp_rank + virtual_index * config.accelerator.pp_size
        for virtual_index in range(config.accelerator.pp_vpp)
    )
300
301
302
303
304
305
306
307
308
        for virtual_index in range(config.accelerator.pp_vpp)
    )
    actual_indices = tuple(chunk.stage_index for chunk in chunks)
    if actual_indices != expected_indices:
        raise ValueError(
            "Dry-run pipeline builder returned unexpected stage indices: "
            f"expected {expected_indices}, got {actual_indices}"
        )
    return chunks
314
315
316
317
318
319
320
321
322
323
324
325
        planner: Any,
        validate_placement: bool,
) -> _DryRunStageShardingResult:
    """Apply the existing TP/CP planner to one already-split stage."""
    mesh = setup.mesh_context
    if mesh.tp_size <= 1 and mesh.cp_size <= 1:
        return _DryRunStageShardingResult(chunk.module, None, None)
    plan = planner.plan(
        chunk.module,
        mesh.device_mesh,
        tp_size=mesh.tp_size,
        cp_size=mesh.cp_size,
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
        ep_size=1,
        sequence_parallel=False,
        loss_parallel=False,
    )
    model, source_shard_info = apply_sharding_plan(
        chunk.module, plan, mesh, validate_mode=validate_placement
    )
    return _DryRunStageShardingResult(model, source_shard_info, plan)


def _active_plan_mesh(mesh: DeviceMesh, plan: ShardingPlan) -> DeviceMesh:
    """Return the stage submesh matching the planner's active dimensions."""
    active_names = tuple(plan.mesh_dim_names)
    if not active_names or tuple(mesh.mesh_dim_names or ()) == active_names:
        return mesh
    selector: str | tuple[str, ...] = active_names[0]
    if len(active_names) > 1:
        selector = active_names
    return mesh[selector]


def _resolve_boundary_metadata(
        leaves: tuple[DryRunBoundaryLeaf, ...],
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
        contract_field: str,
) -> list[list[Any]]:
    """Resolve ordered boundary leaves into PipelineStage metadata."""
    if contract_field not in ("in_dst", "out_dst"):
        raise ValueError(f"Unknown boundary contract field: {contract_field}")
    if mesh.tp_size <= 1 and mesh.cp_size <= 1:
        return [[leaf.global_shape, leaf.dtype, leaf.requires_grad] for leaf in leaves]
    if sharding.plan is None:
        raise ValueError("TP/CP pipeline boundary resolution requires a ShardingPlan")
    metadata = []
    for leaf_index, leaf in enumerate(leaves):
        if not leaf.anchor_fqn:
            raise ValueError(f"TP/CP pipeline boundary leaf {leaf_index} is missing an anchor FQN")
        active_mesh = _active_plan_mesh(mesh.device_mesh, sharding.plan)
        module_spec = sharding.plan.modules.get(leaf.anchor_fqn)
        if module_spec is None:
            raise ValueError(
                f"TP/CP pipeline boundary anchor {leaf.anchor_fqn!r} is absent from the ShardingPlan"
            )
        declared = getattr(module_spec, contract_field)
        if not declared or leaf.tensor_name not in declared:
            raise ValueError(
                f"TP/CP pipeline boundary anchor {leaf.anchor_fqn!r} has no "
                f"{contract_field}[{leaf.tensor_name!r}] declaration"
            )
        placements = tuple(resolve_placements(declared[leaf.tensor_name], sharding.plan.mesh_dim_names))
386
387
388
389
390
391
392
393
394


def _unit_parameters(unit: tuple[nn.Module, ...]) -> set[nn.Parameter]:
    """Return deduplicated trainable parameters owned by one unit."""
    return {
        parameter
        for module in unit
        for parameter in module.parameters()
        if parameter.requires_grad
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
        fsdp_manager: Any,
        fsdp_device_context: Callable[[], ContextManager[Any]],
) -> None:
    """Manually fully-shard declared units while keeping the stage root plain."""
    if fsdp_manager is None:
        return
    trainable_parameters = {parameter for parameter in chunk.module.parameters() if parameter.requires_grad}
    parameters_by_unit = []
    managed_parameters: set[nn.Parameter] = set()
    for unit_index, unit in enumerate(chunk.fsdp_units):
        unit_parameters = _unit_parameters(unit)
        overlap = managed_parameters.intersection(unit_parameters)
        if overlap:
            raise ValueError(
                f"Dry-run pipeline FSDP unit {unit_index} manages {len(overlap)} duplicate parameters"
            )
        managed_parameters.update(unit_parameters)
        parameters_by_unit.append(unit_parameters)
    missing_parameters = trainable_parameters.difference(managed_parameters)
    unexpected_parameters = managed_parameters.difference(trainable_parameters)
    if missing_parameters or unexpected_parameters:
        raise ValueError(
            "Dry-run pipeline FSDP units must manage every trainable parameter exactly once: "
            f"missing={len(missing_parameters)}, unexpected={len(unexpected_parameters)}"
        )

    metadata_by_parameter = _build_source_shard_info_by_param(
        fsdp_manager,
        chunk.module,
        sharding.source_shard_info,
    )
    replicate_parameters = fsdp_manager._resolve_replicate_params(chunk.module)
    fsdp_mesh = fsdp_manager._build_fsdp_actual_mesh()
    fsdp_kwargs, _ = fsdp_manager._build_fully_shard_kwargs(fsdp_mesh)
    default_source_info = (
        _get_default_source_shard_info(fsdp_manager)
        if metadata_by_parameter is not None
        else None
    )
    wrapped_modules = []
    with fsdp_device_context():
        for unit, unit_parameters in zip(chunk.fsdp_units, parameters_by_unit):
            source_infos = None
            if metadata_by_parameter is not None:
                source_infos = {
                    parameter: metadata_by_parameter.get(parameter, default_source_info)
                    for parameter in unit_parameters
                }
            source_infos_for_shard = _source_infos_for_fully_shard(source_infos)
            unit_replicate_parameters = None
            if replicate_parameters is not None:
                unit_replicate_parameters = unit_parameters.intersection(replicate_parameters)
            fully_shard(
                unit[0] if len(unit) == 1 else list(unit),
                source_shard_infos=source_infos_for_shard,
                replicate_params=unit_replicate_parameters,
                **fsdp_kwargs,
455
456
457
458
459
460
461
462
463
464
465
466
467
468
                source_shard_infos=source_infos_for_shard,
                replicate_params=unit_replicate_parameters,
                **fsdp_kwargs,
            )
            control_module = unit[0]
            fsdp_manager._configure_source_layout_gradient_scaling(control_module, source_infos)
            wrapped_modules.append(control_module)
    fsdp_manager._configure_prefetch(wrapped_modules)
    if isinstance(chunk.module, HSDPModule):
        raise RuntimeError("Dry-run pipeline stage execution root must remain unwrapped")


def prepare_pipeline_chunks(
        chunks: tuple[DryRunPipelineChunk, ...],
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
        validate_placement: bool,
        fsdp_device_context: Callable[[], ContextManager[Any]] = nullcontext,
) -> tuple[_PreparedPipelineChunk, ...]:
    """Apply normal-builder stage-local TP/FSDP and resolve wire metadata."""
    prepared = []
    for chunk in chunks:
        sharding = _plan_and_apply_pipeline_chunk(
            chunk, setup, planner, validate_placement
        )
        if sharding.model is not chunk.module:
            raise RuntimeError("Stage-local sharding must preserve the pipeline module identity")
        _apply_pipeline_fsdp(chunk, sharding, fsdp_manager, fsdp_device_context)
        prepared.append(
            _PreparedPipelineChunk(
                chunk,
                _resolve_boundary_metadata(chunk.input_boundary, sharding, setup.mesh_context, "in_dst"),
                _resolve_boundary_metadata(chunk.output_boundary, sharding, setup.mesh_context, "out_dst"),
487
488
489
490
491
492
493
494
495
                _resolve_boundary_metadata(chunk.input_boundary, sharding, setup.mesh_context, "in_dst"),
                _resolve_boundary_metadata(chunk.output_boundary, sharding, setup.mesh_context, "out_dst"),
            )
        )
    return tuple(prepared)


__all__ = [
    "_DryRunParallelContext",