How to get memory_mask for nn.TransformerDecoder

I don’t think so. You don’t need to use memory_mask unless you want to prevent the decoder from attending some tokens in the input sequence, and the original Transformer didn’t use it in the first place because the decoder should be aware of the entire input sequence for any token in the output sequence. The same thing can be said to the input sequence (i.e., src_mask.)

In the PyTorch language, the original Transformer settings are src_mask=None and memory_mask=None, and for tgt_mask=generate_square_subsequent_mask(T).

Again, memory_mask is used only when you don’t want to let the decoder attend certain tokens in the input sequence. That is why the input shape is (S, T) (where S is input sequence lenegth and T is output sequence length.)

If you still want to create a mask, let’s say, so that the decoder does not attend the future positions in the encoder, I’d consider using torch.concat() to create such a mask. For example,

# if S > T

>>> torch.cat([model.generate_square_subsequent_mask(T),
               torch.zeros(T - S, T)], 0)
			   
tensor([[0., -inf, -inf, -inf, -inf],
        [0., 0., -inf, -inf, -inf],
        [0., 0., 0., -inf, -inf],
        [0., 0., 0., 0., -inf],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],				
        [0., 0., 0., 0., 0.]])

# if T > S
>>> torch.cat([model.generate_square_subsequent_mask(S),
               torch.zeros(S, T - S)], 1)
			   
tensor([[0., -inf, -inf, -inf, -inf, 0., 0.],
        [0., 0., -inf, -inf, -inf, 0., 0.],
        [0., 0., 0., -inf, -inf, 0., 0.],
        [0., 0., 0., 0., -inf, 0., 0.],
        [0., 0., 0., 0., 0., 0., 0.]])		

Hope it helps.