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)
|