File size: 4,926 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch

import torch.nn as nn
import torch.nn.functional as F
import torch.distributed as dist
import os

from diffulex.distributed.parallel_state import fetch_parallel_state
from diffulex.vllm_compat import get_vllm_tp_group


LM_HEAD_FP32 = os.environ.get("DIFFULEX_LM_HEAD_FP32", "0") == "1"
LM_HEAD_FP32_GATHER = LM_HEAD_FP32 or os.environ.get("DIFFULEX_LM_HEAD_FP32_GATHER", "0") == "1"


def _tp_all_reduce(x: torch.Tensor, group) -> torch.Tensor:
    vllm_tp_group = get_vllm_tp_group()
    if vllm_tp_group is not None:
        try:
            return vllm_tp_group.all_reduce(x)
        except Exception:
            pass
    dist.all_reduce(x, group=group)
    return x


def _tp_gather_to_rank0(x: torch.Tensor, group, tp_size: int, tp_rank: int) -> torch.Tensor | None:
    vllm_tp_group = get_vllm_tp_group()
    if vllm_tp_group is not None:
        try:
            return vllm_tp_group.gather(x, dst=0, dim=-1)
        except Exception:
            pass

    gathered = [torch.empty_like(x) for _ in range(tp_size)]
    dist.all_gather(gathered, x, group=group)
    return torch.cat(gathered, -1) if tp_rank == 0 else None


class VocabParallelEmbedding(nn.Module):
    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
    ):
        super().__init__()
        parallel_state = fetch_parallel_state()
        self.tp_rank = parallel_state.get_tp_rank()
        self.tp_size = parallel_state.get_tp_world_size()
        self.tp_group = parallel_state.get_tp_group()
        assert num_embeddings % self.tp_size == 0
        self.num_embeddings = num_embeddings
        self.num_embeddings_per_partition = self.num_embeddings // self.tp_size
        self.vocab_start_idx = self.num_embeddings_per_partition * self.tp_rank
        self.vocab_end_idx = self.vocab_start_idx + self.num_embeddings_per_partition
        self.weight = nn.Parameter(torch.empty(self.num_embeddings_per_partition, embedding_dim))
        self.weight.weight_loader = self.weight_loader

    def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
        param_data = param.data
        shard_size = param_data.size(0)
        start_idx = self.tp_rank * shard_size
        loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
        assert param_data.size() == loaded_weight.size()
        param_data.copy_(loaded_weight)

    def forward(self, x: torch.Tensor):
        if self.tp_size > 1:
            mask = (x >= self.vocab_start_idx) & (x < self.vocab_end_idx)
            x = mask * (x - self.vocab_start_idx)
        y = F.embedding(x, self.weight)
        if self.tp_size > 1:
            y = mask.unsqueeze(1) * y
            y = _tp_all_reduce(y, self.tp_group)
        return y


class ParallelLMHead(VocabParallelEmbedding):
    def __init__(
        self,
        num_embeddings: int,
        embedding_dim: int,
        bias: bool = False,
    ):
        super().__init__(num_embeddings, embedding_dim)
        if bias:
            self.bias = nn.Parameter(torch.empty(self.num_embeddings_per_partition))
            self.bias.weight_loader = self.weight_loader
        else:
            self.register_parameter("bias", None)
        self._local_logits_workspace: torch.Tensor | None = None

    def _get_local_logits_workspace(self, x: torch.Tensor) -> torch.Tensor:
        shape = (*x.shape[:-1], self.num_embeddings_per_partition)
        workspace = self._local_logits_workspace
        if (
            workspace is None
            or tuple(workspace.shape) != tuple(shape)
            or workspace.device != x.device
            or workspace.dtype != x.dtype
        ):
            workspace = torch.empty(shape, device=x.device, dtype=x.dtype)
            self._local_logits_workspace = workspace
        return workspace

    def _linear_into_workspace(self, x: torch.Tensor) -> torch.Tensor:
        if LM_HEAD_FP32:
            return F.linear(
                x.to(torch.float32),
                self.weight.to(torch.float32),
                self.bias.to(torch.float32) if self.bias is not None else None,
            ).to(x.dtype)

        if x.dim() != 2:
            return F.linear(x, self.weight, self.bias)

        logits = self._get_local_logits_workspace(x)
        torch.mm(x, self.weight.t(), out=logits)
        if self.bias is not None:
            logits.add_(self.bias)
        return logits

    def forward(self, x: torch.Tensor):
        logits = self._linear_into_workspace(x)
        if self.tp_size > 1:
            if LM_HEAD_FP32_GATHER:
                logits_dtype = logits.dtype
                logits = _tp_gather_to_rank0(logits.to(torch.float32), self.tp_group, self.tp_size, self.tp_rank)
                logits = logits.to(logits_dtype) if logits is not None else None
            else:
                logits = _tp_gather_to_rank0(logits, self.tp_group, self.tp_size, self.tp_rank)
        return logits