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