File size: 7,616 Bytes
e19323e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
144
145
146
147
148
149
150
151
152
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
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
# Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.
# For a list of all contributors, visit:
#   https://github.com/fla-org/flash-linear-attention/graphs/contributors

"""
Context Parallel support for Token Shift.

Token shift has a 1-token dependency on previous tokens:
    y[t] = x[t-1] - x[t]  (for t > 0)
    y[0] = cache - x[0]   (cache is the last token from previous rank)

In CP mode, non-first ranks need the last token from the previous rank as cache.
Backward: non-last ranks need to send the last token's gradient to previous rank.
"""

import torch
import torch.distributed as dist

from ..modules.token_shift import token_shift_bwd, token_shift_fwd
from ..ops.cp import FLACPContext, conv_cp_send_recv_bwd, conv_cp_send_recv_fwd


class TokenShiftCPFunction(torch.autograd.Function):
    """
    Context Parallel version of TokenShift.

    Forward:
        1. Get last token from previous rank to construct cache
        2. Call token_shift_fwd with cache

    Backward:
        1. Call token_shift_bwd to get dx
        2. Sync communication: add next rank's first token gradient to current rank's last token
    """

    @staticmethod
    def _prepare_cache_for_cp(
        x: torch.Tensor,
        cu_seqlens: torch.Tensor | None,
        context: FLACPContext,
        group: dist.ProcessGroup | None,
    ) -> tuple[torch.Tensor | None, int]:
        """Prepare cache for CP forward pass by communicating with previous rank.

        Args:
            x: Input tensor of shape [1, T, D]
            cu_seqlens: Cumulative sequence lengths
            context: CP context
            group: Process group for communication

        Returns:
            cache: Cache tensor of shape [N, D] or None
            pre_num_tokens: Number of tokens from previous rank for the first sequence
        """
        if group is None:
            return None, 0

        D = x.shape[-1]
        cache = None
        pre_num_tokens = 0

        if not context.is_first_rank:
            # Non-first rank: need cache from previous rank
            assert x.dim() == 3 and x.shape[0] == 1, f"CP requires [1, T, D], got {x.shape}"
            x_2d = x.squeeze(0)  # [T, D]
            last_token = x_2d[-1:].contiguous()  # [1, D]
            prev_last_token = conv_cp_send_recv_fwd(last_token, group)  # [1, D]

            # For varlen: only the first sequence needs cache from prev rank
            N = len(cu_seqlens) - 1 if cu_seqlens is not None else 1
            cache = torch.zeros(N, D, device=x.device, dtype=x.dtype)

            # pre_num_conv_tokens tells us how many tokens from prev rank
            # belong to the first sequence on this rank
            pre_num_tokens = getattr(context, 'pre_num_conv_tokens', 0)
            if pre_num_tokens > 0:
                # The prev rank's last token is used as cache for first sequence
                cache[0] = prev_last_token[0]
        else:
            # First rank: participate in send but don't use received data
            x_2d = x.squeeze(0)
            last_token = x_2d[-1:].contiguous()
            _ = conv_cp_send_recv_fwd(last_token, group)

        return cache, pre_num_tokens

    @staticmethod
    def _correct_dx_for_cp(
        dx: torch.Tensor,
        grad_cache: torch.Tensor | None,
        group: dist.ProcessGroup | None,
        is_first_rank: bool,
        pre_num_tokens: int = 0,
    ) -> None:
        """Correct dx gradients for CP backward pass.

        Args:
            dx: Gradient tensor to be corrected, shape [1, T, D]
            grad_cache: Gradient w.r.t. cache, shape [N, D] or None
            group: Process group
            is_first_rank: Whether this is the first rank
            pre_num_tokens: Number of tokens from previous rank for first sequence
        """
        if group is None:
            return

        D = dx.shape[-1]

        # Prepare gradient to send to previous rank
        if grad_cache is not None and pre_num_tokens > 0:
            # Only first sequence's cache gradient is relevant
            d_cache = grad_cache[0:1]  # [1, D]
        else:
            d_cache = torch.zeros(1, D, device=dx.device, dtype=dx.dtype)

        # Send to previous rank, receive from next rank
        recv_grad = conv_cp_send_recv_bwd(d_cache, group)  # [1, D]

        # Add received gradient to current rank's last token
        dx[0, -1, :].add_(recv_grad[0])

    @staticmethod
    def forward(
        ctx,
        x: torch.Tensor,
        cu_seqlens: torch.Tensor | None,
        chunk_indices: torch.Tensor | None,
        cp_context: FLACPContext | None,
    ):
        if cp_context is None:
            raise ValueError("cp_context must be provided for TokenShiftCPFunction")

        cu_seqlens = cp_context.cu_seqlens
        group = cp_context.group

        # Prepare cache for CP
        cache, pre_num_tokens = TokenShiftCPFunction._prepare_cache_for_cp(
            x=x,
            cu_seqlens=cu_seqlens,
            context=cp_context,
            group=group,
        )

        # Save for backward
        ctx.cu_seqlens = cu_seqlens
        ctx.chunk_indices = chunk_indices
        ctx.group = group
        ctx.has_cache = cache is not None
        ctx.is_first_rank = cp_context.is_first_rank
        ctx.pre_num_tokens = pre_num_tokens

        # Call original forward
        y, N, T, use_short_kernel, cache_out = token_shift_fwd(
            x=x,
            cu_seqlens=cu_seqlens,
            cache=cache,
            output_cache=True,
            chunk_indices=chunk_indices,
        )

        ctx.N = N
        ctx.T = T
        ctx.use_short_kernel = use_short_kernel

        return y

    @staticmethod
    def backward(ctx, dy: torch.Tensor):
        group = ctx.group

        # Prepare dcache for backward
        # For CP: non-last rank needs to receive gradient from next rank
        # This is handled in _correct_dx_for_cp after computing dx
        dcache = None  # Will be computed by token_shift_bwd

        # Call original backward
        dx, grad_cache = token_shift_bwd(
            dy=dy,
            N=ctx.N,
            T=ctx.T,
            dcache=dcache,
            cu_seqlens=ctx.cu_seqlens,
            use_short_kernel=ctx.use_short_kernel,
            has_init_cache=ctx.has_cache,
            chunk_indices=ctx.chunk_indices,
        )

        # Correct dx gradients for CP
        TokenShiftCPFunction._correct_dx_for_cp(
            dx=dx,
            grad_cache=grad_cache,
            group=group,
            is_first_rank=ctx.is_first_rank,
            pre_num_tokens=ctx.pre_num_tokens,
        )

        return dx, None, None, None


@torch.compiler.disable
def token_shift_cp(
    x: torch.Tensor,
    cp_context: FLACPContext,
    cu_seqlens: torch.Tensor | None = None,
    chunk_indices: torch.Tensor | None = None,
):
    """
    Context Parallel version of token_shift.

    Args:
        x: Input tensor of shape [1, T, D]
        cp_context: CP context (required for CP mode)
        cu_seqlens: Cumulative sequence lengths
        chunk_indices: Chunk indices for variable-length sequences

    Returns:
        output: Tensor of shape [1, T, D] after applying token-shift
    """
    if cp_context is None:
        raise ValueError("cp_context must be provided for token_shift_cp")

    assert cp_context.cu_seqlens is not None, "cu_seqlens must be provided for token_shift_cp"

    return TokenShiftCPFunction.apply(
        x, cu_seqlens, chunk_indices, cp_context
    )