File size: 4,536 Bytes
d91766b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 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 81 82 83 84 85 86 87 88 89 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 | from __future__ import annotations
from dataclasses import dataclass
from enum import Enum, auto
import torch
from diffulex.moe.topk.output import TopKOutput
@dataclass(frozen=True)
class RouterMetadata:
router_logits: torch.Tensor
topk_ids: torch.Tensor
topk_weights: torch.Tensor
@classmethod
def empty(
cls,
hidden_states: torch.Tensor,
*,
num_experts: int,
top_k: int,
) -> "RouterMetadata":
"""Empty router output for idle/local-empty ranks in EP collectives."""
return cls(
router_logits=hidden_states.new_empty((0, num_experts)),
topk_ids=torch.full(
(0, top_k),
-1,
device=hidden_states.device,
dtype=torch.int32,
),
topk_weights=torch.empty(
(0, top_k),
device=hidden_states.device,
dtype=hidden_states.dtype,
),
)
@classmethod
def from_topk_output(cls, topk_output: TopKOutput) -> "RouterMetadata":
return cls(
router_logits=topk_output.router_logits,
topk_ids=topk_output.ids,
topk_weights=topk_output.weights,
)
class DispatcherStage(Enum):
INITIAL = auto()
AFTER_DISPATCH_A = auto()
AFTER_DISPATCH_B = auto()
AFTER_COMBINE_A = auto()
@dataclass(frozen=True)
class DispatchMetadata:
num_tokens: int
hidden_size: int
dtype: torch.dtype
device: torch.device
send_splits: list[int]
recv_splits: list[int]
recv_hidden_states: torch.Tensor
recv_local_expert: torch.Tensor
recv_token_indices: torch.Tensor
recv_weights: torch.Tensor
total_recv_slots: int
active_dispatch: bool = True
num_local_tokens: int | None = None
local_token_indices: torch.Tensor | None = None
@dataclass(frozen=True)
class ExpertExecutionMetadata:
packed_token_ids: torch.Tensor
packed_local_expert_ids: torch.Tensor
packed_weights: torch.Tensor
num_slots: int
seg_indptr: torch.Tensor | None = None
num_recv_tokens_per_expert: torch.Tensor | None = None
sorted_slot_ids: torch.Tensor | None = None
expert_block_ids: torch.Tensor | None = None
num_tokens_post_padded: int | None = None
disable_aligned_metadata: bool = False
@dataclass(frozen=True)
class DeepEPDispatchMetadata(DispatchMetadata):
src2dst: torch.Tensor | None = None
reorder_indices: torch.Tensor | None = None
reordered_token_indices: torch.Tensor | None = None
reordered_local_expert_ids: torch.Tensor | None = None
seg_indptr: torch.Tensor | None = None
num_recv_tokens_per_expert: torch.Tensor | None = None
native_handle: object | None = None
native_recv_num_tokens: int | None = None
native_recv_topk_ids: torch.Tensor | None = None
native_recv_topk_weights: torch.Tensor | None = None
low_latency: bool = False
low_latency_handle: object | None = None
low_latency_topk_ids: torch.Tensor | None = None
low_latency_topk_weights: torch.Tensor | None = None
low_latency_recv_count: torch.Tensor | None = None
low_latency_capacity: int | None = None
def to_expert_execution_metadata(self) -> ExpertExecutionMetadata:
recv_local_expert = (
self.reordered_local_expert_ids
if self.reordered_local_expert_ids is not None
else self.recv_local_expert
)
if recv_local_expert is None:
raise ValueError("DeepEPDispatchMetadata is missing recv_local_expert information.")
recv_weights = self.recv_weights
if recv_weights is None:
raise ValueError("DeepEPDispatchMetadata is missing recv_weights.")
num_slots = int(recv_local_expert.numel())
packed_token_ids = torch.arange(
num_slots,
device=recv_local_expert.device,
dtype=torch.int32,
)
return ExpertExecutionMetadata(
packed_token_ids=packed_token_ids,
packed_local_expert_ids=recv_local_expert.to(torch.int32).contiguous(),
packed_weights=recv_weights.contiguous(),
num_slots=num_slots,
seg_indptr=self.seg_indptr,
num_recv_tokens_per_expert=self.num_recv_tokens_per_expert,
disable_aligned_metadata=self.low_latency,
)
__all__ = [
"DispatchMetadata",
"DeepEPDispatchMetadata",
"DispatcherStage",
"ExpertExecutionMetadata",
"RouterMetadata",
]
|