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