src_len, tgt_len = torch.tensor([3, 5]), torch.tensor([3, 4])
S, T_dec = int(src_len.max()), int(tgt_len.max())
src_valid = (torch.arange(S)[None, :] < src_len[:, None])[:, None, :]
tgt_valid = (torch.arange(T_dec)[None, :] < tgt_len[:, None])[:, None, :]
causal = (torch.arange(T_dec)[None, :]
<= torch.arange(T_dec)[:, None])[None] # (1, T, T)
enc_self = src_valid.expand(-1, S, -1) # (B, S, S): source padding
dec_self = causal & tgt_valid # (B, T, T): causal AND padding
cross = src_valid.expand(-1, T_dec, -1) # (B, T, S): source padding
for name, m in (('encoder self', enc_self), ('decoder self', dec_self),
('cross', cross)):
print(f'{name}-attention mask, sequence 0 (1 = may attend):')
print(m[0].int().numpy())