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:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user