From bfa372d9c3f5640dd0043aaf0e25a9e16ce59055 Mon Sep 17 00:00:00 2001 From: Mateo Saldain Date: Tue, 28 Jul 2026 00:57:40 -0300 Subject: [PATCH] La z-loss ya no duplica la matriz de logits MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit La verificación previa contra la 2060 real hizo fallar el pre-entrenamiento por falta de memoria: 11,19 de 11,56 GB en uso. La causa es que la z-loss indexaba los logits con una máscara booleana antes del logsumexp, y eso copia la matriz completa. Con vocabulario de 32k y lotes de 16k tokens son 2 GB de más. Ahora el logsumexp se calcula sobre todas las filas y el filtro se aplica al vector resultante, que tiene una entrada por token en vez de una por token y clase. El resultado numérico es idéntico. Encontrado por la verificación previa, no por los tests: los tests corren en CPU con modelos diminutos, donde estos 2 GB son unos pocos megabytes. Co-Authored-By: Claude Opus 5 --- enlace/model/transformer.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) 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