| 8 | |
| 9 | |
| 10 | def create_sparse_mask(document_lens, split_lens, attn_modes, device): |
| 11 | def causal_mask(b, h, q_idx, kv_idx): |
| 12 | return q_idx >= kv_idx |
| 13 | |
| 14 | def full_and_noise_mask(b, h, q_idx, kv_idx): |
| 15 | return (full_and_noise_seq_id[q_idx] == full_and_noise_seq_id[kv_idx]) & (full_and_noise_seq_id[q_idx] >= 0) |
| 16 | |
| 17 | def remove_noise_mask(b, h, q_idx, kv_idx): |
| 18 | return (~((noise_seq_id[kv_idx] >= 0) & (noise_seq_id[q_idx] != noise_seq_id[kv_idx]))) |
| 19 | |
| 20 | def sample_mask(b, h, q_idx, kv_idx): |
| 21 | return document_id[q_idx] == document_id[kv_idx] |
| 22 | |
| 23 | full_and_noise_tmp = [] |
| 24 | noise_tmp = [] |
| 25 | |
| 26 | for i, (length, model) in enumerate(zip(split_lens, attn_modes)): |
| 27 | value = i if model in ['full', 'noise'] else -1 |
| 28 | full_and_noise_tmp.extend([value] * length) |
| 29 | value_noise = i if model == 'noise' else -1 |
| 30 | noise_tmp.extend([value_noise] * length) |
| 31 | |
| 32 | full_and_noise_seq_id = torch.Tensor(full_and_noise_tmp).to(device) |
| 33 | noise_seq_id = torch.Tensor(noise_tmp).to(device) |
| 34 | |
| 35 | document_id = torch.cat([torch.full((l,), i) for i, l in enumerate(document_lens, start=1)]).to(device) |
| 36 | |
| 37 | return and_masks(or_masks(causal_mask, full_and_noise_mask), remove_noise_mask, sample_mask) |
| 38 | |
| 39 | |
| 40 | def patchify(image, patch_size): |