INNER CODE UNIT · Python

lists

lucidrains/x-transformers · train_discover_reversal.py:39

def lists(n, lengths, device, num_fillers):
    max_len = int(lengths.max())
    tokens = torch.randint(0, num_fillers, (n, max_len), device = device)
    valid = einx.less('n, b -> b n', torch.arange(max_len, device = device), lengths)
    return tokens.masked_fill(~valid, num_fillers + 1)

def reversal_targets(tokens, lengths, device, num_fillers):
    batch, seq_len = tokens.shape
    pad_id = num_fillers + 1
    positions = torch.arange(seq_len, device = device)
    valid = einx.less('n, b -> b n', positions, lengths)
    indices = einx.subtract('b, n -> b n', lengths - 1, positions).clamp(min = 0)
    targets = tokens[torch.arange(batch, device = device)[:, None], indices]
    return targets.masked_fill(~valid, pad_id)

def main(
    *,
    num_fillers: int = 26,

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…