La z-loss ya no duplica la matriz de logits

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 <noreply@anthropic.com>
This commit is contained in:
2026-07-28 00:57:40 -03:00
parent 12bdc983e6
commit bfa372d9c3
+10 -2
View File
@@ -217,11 +217,19 @@ class Transformer(nn.Module):
# z-loss: penaliza que los logits crezcan en magnitud absoluta. Es lo # z-loss: penaliza que los logits crezcan en magnitud absoluta. Es lo
# que impide que el softmax sature y desborde entrenando en fp16. # 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: if self.cfg.z_loss_weight > 0.0:
z = torch.logsumexp(flat_logits, dim=-1)
valid = flat_targets != -100 valid = flat_targets != -100
if valid.any(): if valid.all():
z = torch.logsumexp(flat_logits[valid], dim=-1)
loss = loss + self.cfg.z_loss_weight * z.pow(2).mean() 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 return logits, loss