block_size = attn_scores.shape[0] mask_simple = torch.tril(torch.ones(block_size, block_size)) print(mask_simple)