Spaces:
Running
Running
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Any | |
| from .duration import op_type_value | |
| COMPUTE_OP_TYPES = frozenset({"F", "B", "VF", "VB", "GF", "GB"}) | |
| class PlanOp: | |
| id: str | |
| rank: int | |
| index: int | |
| original_index: int | |
| op_type: str | |
| microbatch_id: int | None | |
| chunk_id: int | None | |
| base_duration: float | |
| label: str | |
| def is_forward(self) -> bool: | |
| return self.op_type == "F" | |
| def is_backward(self) -> bool: | |
| return self.op_type == "B" | |
| class Plan: | |
| ops_by_rank: dict[int, list[PlanOp]] | |
| pp_size: int | |
| num_microbatches: int | |
| scheduler_name: str | |
| vpp_size: int = 1 | |
| pipeline_layout: dict[str, int] | None = None | |
| def ops(self) -> list[PlanOp]: | |
| return [ | |
| op | |
| for rank in sorted(self.ops_by_rank) | |
| for op in self.ops_by_rank[rank] | |
| ] | |
| def from_scheduler(cls, scheduler: Any, *, include_op_types=None) -> "Plan": | |
| if not getattr(scheduler, "schedules", None): | |
| scheduler.generate_schedule() | |
| allowed = frozenset(include_op_types or COMPUTE_OP_TYPES) | |
| ops_by_rank: dict[int, list[PlanOp]] = {} | |
| pp_size = int(getattr(scheduler, "pp_size")) | |
| for rank in range(pp_size): | |
| rank_ops: list[PlanOp] = [] | |
| for original_index, raw_op in enumerate(scheduler.schedules.get(rank, [])): | |
| op_value = op_type_value(getattr(raw_op, "op_type", "")) | |
| if op_value not in allowed: | |
| continue | |
| rank_index = len(rank_ops) | |
| microbatch_id = getattr(raw_op, "microbatch_id", None) | |
| chunk_id = getattr(raw_op, "chunk_id", None) | |
| op_id = _make_op_id(rank, rank_index, op_value, microbatch_id, chunk_id) | |
| rank_ops.append( | |
| PlanOp( | |
| id=op_id, | |
| rank=rank, | |
| index=rank_index, | |
| original_index=original_index, | |
| op_type=op_value, | |
| microbatch_id=microbatch_id, | |
| chunk_id=chunk_id, | |
| base_duration=float(getattr(raw_op, "duration", 0.0)), | |
| label=str(raw_op), | |
| ) | |
| ) | |
| ops_by_rank[rank] = rank_ops | |
| return cls( | |
| ops_by_rank=ops_by_rank, | |
| pp_size=pp_size, | |
| num_microbatches=int(getattr(scheduler, "num_microbatches")), | |
| scheduler_name=scheduler.__class__.__name__, | |
| vpp_size=int(getattr(scheduler, "vpp_size", 1)), | |
| pipeline_layout=getattr(scheduler, "pipeline_layout", None), | |
| ) | |
| def _make_op_id( | |
| rank: int, | |
| index: int, | |
| op_type: str, | |
| microbatch_id: int | None, | |
| chunk_id: int | None, | |
| ) -> str: | |
| mb = "none" if microbatch_id is None else str(microbatch_id) | |
| chunk = "none" if chunk_id is None else str(chunk_id) | |
| return f"r{rank}:i{index}:{op_type}:mb{mb}:c{chunk}" | |