diff --git a/enlace/model/transformer.py b/enlace/model/transformer.py index 1f7f9c0..864c769 100644 --- a/enlace/model/transformer.py +++ b/enlace/model/transformer.py @@ -217,11 +217,19 @@ class Transformer(nn.Module): # z-loss: penaliza que los logits crezcan en magnitud absoluta. Es lo # que impide que el softmax sature y desborde entrenando en fp16. + # + # El logsumexp se calcula sobre todas las filas y recién después se + # filtra. Indexar antes (`flat_logits[valid]`) copiaría la matriz + # entera de logits — con vocabulario de 32k y lotes de 16k tokens eso + # son 2 GB de más, suficiente para que la corrida no entre en una 2060. + # Filtrar el vector resultante, de una entrada por token, no cuesta nada. if self.cfg.z_loss_weight > 0.0: + z = torch.logsumexp(flat_logits, dim=-1) valid = flat_targets != -100 - if valid.any(): - z = torch.logsumexp(flat_logits[valid], dim=-1) + if valid.all(): loss = loss + self.cfg.z_loss_weight * z.pow(2).mean() + elif valid.any(): + loss = loss + self.cfg.z_loss_weight * z[valid].pow(2).mean() return logits, loss