Skip to content

GroupSwapLayer produces NaN when all tokens land in the same partition group #1

Description

@deuterium1729

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)

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions