ObsidianSmall-Base / runtime /litgpt /multiscreen_loss.py
Metris's picture
Upload 78 files
236083b verified
Raw
History Blame Contribute Delete
3.01 kB
"""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)