MCPcopy Create free account
hub / github.com/InternRobotics/G2VLM / create_sparse_mask

Function create_sparse_mask

data/data_utils.py:10–37  ·  view source on GitHub ↗
(document_lens, split_lens, attn_modes, device)

Source from the content-addressed store, hash-verified

8
9
10def 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
40def patchify(image, patch_size):

Callers 1

forwardMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected