Implement Transformer-XL Recurrence
mainTo implement Transformer-XL recurrence:
- In
TransformerWrapper, setmax_mem_len(e.g., 2048). - In the
DecoderorEncoder, setrel_pos_bias = Trueorrotary_pos_emb = True. - Use
return_mems = Trueduring the forward pass to get memories. - Pass the retrieved memories back into the next iteration using the
memskeyword.
import torch
from x_transformers import TransformerWrapper, Decoder
model_xl = TransformerWrapper(
num_tokens = 20000,
max_seq_len = 512,
max_mem_len = 2048,
attn_layers = Decoder(
dim = 512,
depth = 6,
heads = 8,
rel_pos_bias = True
)
)
seg1 = torch.randint(0, 20000, (1, 512))
seg2 = torch.randint(0, 20000, (1, 512))
logits1, mems1 = model_xl(seg1, return_mems = True)
logits2, mems2 = model_xl(seg2, mems = mems1, return_mems = True)