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