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",
]