src_len, tgt_len = jnp.array([3, 5]), jnp.array([3, 4])
S, T_dec = int(src_len.max()), int(tgt_len.max())
src_valid = (jnp.arange(S)[None, :] < src_len[:, None])[:, None, :]
tgt_valid = (jnp.arange(T_dec)[None, :] < tgt_len[:, None])[:, None, :]
causal = (jnp.arange(T_dec)[None, :]
<= jnp.arange(T_dec)[:, None])[None] # (1, T, T)
B = len(src_len)
enc_self = jnp.broadcast_to(src_valid, (B, S, S)) # source padding
dec_self = causal & tgt_valid # (B, T, T): causal AND padding
cross = jnp.broadcast_to(src_valid, (B, T_dec, 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].astype(int))