INNER CODE UNIT · Python

exists

lucidrains/x-transformers · train_discover_reversal.py:33

def exists(v):
    return v is not None

def to_letters(tokens, num_fillers = 26):
    return ''.join(chr(65 + v) for v in tokens if v < num_fillers)

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)

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…