INNER CODE UNIT · Python

reversal_targets

lucidrains/x-transformers · train_discover_reversal.py:45

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,
    len_max: int = 3,
    batch_size: int = 128,
    train_steps: int = 12000,
    dim: int = 192,
    depth: int = 4,
    heads: int = 6,

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…