Summary
Pre-Training on the text-pretrain branch with the provided settings (small-encoder-decoder, partition-mdlm, lm1b) produces NaN model outputs from the very first forward pass (including sanity check, before any gradient updates). Thus all results of the training report NaN including val/ppl, val/nll, etc. I think cause is an all-False cross-attention mask in GroupSwapLayer when alpha_t is near 0 or 1.
Environment
Reproduces with the exact settings from the repo using a single a100 gpu:
- model=small-encoder-decoder
- algo=partition-mdlm, data=lm1b
- trainer.precision="32-true"
- trainer.num_nodes=1
- trainer.devices=1
- loader.batch_size=256
Possible issues: partition sampling
_q_xt_partition assigns each token independently to group 0 (prob alpha_t) or group 1 (prob 1 - alpha_t):
# algo.py
def _q_xt_partition(self, x, alpha_t):
group_idxs = torch.rand(*x.shape, device=x.device) < 1 - alpha_t
return group_idxs.to(int)
With stratified antithetic sampling (_sample_t), t is evenly spread across [0.001, 1]. The first elements of each batch get t = 0.001 so alpha_t = 0.999, and the last elements get t = 0.999 so alpha_t = 0.001. At these extremes, the probability that all 128 tokens land in the same group is:
P(all group 0 | alpha_t = 0.999) = 0.999^128 = 0.878
P(all group 1 | alpha_t = 0.001) = 0.999^128 = 0.878
For the first few and last few sequences in every batch, group_idxs is a constant vector of all 0s or all 1s
Possible issues: All-False cross-attention mask
Decoder.forward() (and GroupSwapLayer.forward()) build the cross-attention mask as:
# models/encoder_decoder.py
def make_group_cross_attn_mask(group_idxs):
return group_idxs[:, None, :] != group_idxs[:, :, None]
when all tokens are in the same group, every row of this mask is false because there are no cross-group keys to attend to.
The result is NaN from softmax over all -inf.
apply_masked_mha passes this boolean mask to F.scaled_dot_product_attention. PyTorch converts False to -inf in the attention logits, and softmax([-inf, ..., -inf]) = NaN.
GroupSwapLayer.forward() returns NaN for these sequences. This then become the initial x in Decoder.forward(), and NaN propagates through all 6 decoder layers via the residual connections.
Possible Fix:
In _q_xt_partition, ensure each sequence has at least one token in each group before returning:
def _q_xt_partition(self, x, alpha_t):
group_idxs = torch.rand(*x.shape, device=x.device) < 1 - alpha_t
# Guarantee both groups are non-empty per sequence
all_ones = group_idxs.all(dim=-1) # all in group 1
all_zeros = ~group_idxs.any(dim=-1) # all in group 0
if all_ones.any():
rand_pos = torch.randint(x.shape[1], (all_ones.sum(),), device=x.device)
group_idxs[all_ones.nonzero(as_tuple=True)[0], rand_pos] = False
if all_zeros.any():
rand_pos = torch.randint(x.shape[1], (all_zeros.sum(),), device=x.device)
group_idxs[all_zeros.nonzero(as_tuple=True)[0], rand_pos] = True
return group_idxs.to(int)
Summary
Pre-Training on the text-pretrain branch with the provided settings (small-encoder-decoder, partition-mdlm, lm1b) produces NaN model outputs from the very first forward pass (including sanity check, before any gradient updates). Thus all results of the training report NaN including val/ppl, val/nll, etc. I think cause is an all-False cross-attention mask in GroupSwapLayer when alpha_t is near 0 or 1.
Environment
Reproduces with the exact settings from the repo using a single a100 gpu:
Possible issues: partition sampling
_q_xt_partitionassigns each token independently to group 0 (prob alpha_t) or group 1 (prob 1 - alpha_t):With stratified antithetic sampling (_sample_t), t is evenly spread across [0.001, 1]. The first elements of each batch get t = 0.001 so alpha_t = 0.999, and the last elements get t = 0.999 so alpha_t = 0.001. At these extremes, the probability that all 128 tokens land in the same group is:
P(all group 0 | alpha_t = 0.999) = 0.999^128 = 0.878
P(all group 1 | alpha_t = 0.001) = 0.999^128 = 0.878
For the first few and last few sequences in every batch, group_idxs is a constant vector of all 0s or all 1s
Possible issues: All-False cross-attention mask
Decoder.forward() (and GroupSwapLayer.forward()) build the cross-attention mask as:
when all tokens are in the same group, every row of this mask is false because there are no cross-group keys to attend to.
The result is NaN from softmax over all -inf.
apply_masked_mhapasses this boolean mask toF.scaled_dot_product_attention. PyTorch converts False to -inf in the attention logits, andsoftmax([-inf, ..., -inf]) = NaN.GroupSwapLayer.forward() returns NaN for these sequences. This then become the initial x in Decoder.forward(), and NaN propagates through all 6 decoder layers via the residual connections.
Possible Fix:
In _q_xt_partition, ensure each sequence has at least one token in each group before returning: