File size: 3,009 Bytes
236083b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Memory-bounded tied language-head cross entropy for Multiscreen training."""

from __future__ import annotations

import torch


class _ChunkedLinearCrossEntropy(torch.autograd.Function):
    @staticmethod
    def forward(
        ctx,
        hidden: torch.Tensor,
        weight: torch.Tensor,
        targets: torch.Tensor,
        chunk_tokens: int,
    ) -> torch.Tensor:
        hidden_flat = hidden.reshape(-1, hidden.size(-1))
        targets_flat = targets.reshape(-1)
        valid_count = (targets_flat != -100).sum().clamp_min(1)
        loss_sum = torch.zeros((), device=hidden.device, dtype=torch.float32)
        for start in range(0, hidden_flat.size(0), chunk_tokens):
            end = min(start + chunk_tokens, hidden_flat.size(0))
            chunk_targets = targets_flat[start:end]
            logits = hidden_flat[start:end].float() @ weight.float().t()
            loss_sum += torch.nn.functional.cross_entropy(
                logits,
                chunk_targets,
                ignore_index=-100,
                reduction="sum",
            )
        ctx.save_for_backward(hidden, weight, targets, valid_count)
        ctx.chunk_tokens = chunk_tokens
        return loss_sum / valid_count

    @staticmethod
    def backward(ctx, grad_output: torch.Tensor):
        hidden, weight, targets, valid_count = ctx.saved_tensors
        hidden_flat = hidden.reshape(-1, hidden.size(-1))
        targets_flat = targets.reshape(-1)
        grad_hidden = torch.empty_like(hidden_flat)
        grad_weight = torch.zeros_like(weight, dtype=torch.float32)
        scale = grad_output.float() / valid_count.float()

        for start in range(0, hidden_flat.size(0), ctx.chunk_tokens):
            end = min(start + ctx.chunk_tokens, hidden_flat.size(0))
            hidden_chunk = hidden_flat[start:end].float()
            chunk_targets = targets_flat[start:end]
            logits = hidden_chunk @ weight.float().t()
            probabilities = logits.softmax(dim=-1)
            valid = chunk_targets != -100
            safe_targets = chunk_targets.masked_fill(~valid, 0)
            probabilities[
                torch.arange(end - start, device=hidden.device), safe_targets
            ] -= valid.to(probabilities.dtype)
            probabilities *= valid[:, None].to(probabilities.dtype)
            probabilities *= scale
            grad_hidden[start:end] = (
                probabilities @ weight.float()
            ).to(hidden.dtype)
            grad_weight.add_(probabilities.t() @ hidden_chunk)

        return (
            grad_hidden.view_as(hidden),
            grad_weight.to(weight.dtype),
            None,
            None,
        )


def chunked_linear_cross_entropy(
    hidden: torch.Tensor,
    weight: torch.Tensor,
    targets: torch.Tensor,
    chunk_tokens: int = 2048,
) -> torch.Tensor:
    if chunk_tokens <= 0:
        raise ValueError("chunk_tokens must be positive")
    return _ChunkedLinearCrossEntropy.apply(hidden, weight, targets, chunk_tokens)