| import torch |
|
|
| def make_basic_block_attention( |
| N: int, |
| start_pos: int, |
| block_size: int, |
| ) -> torch.Tensor: |
| B = 1 |
| L0 = start_pos |
| L1 = (N - L0) // 2 |
| assert L0 + 2 * L1 == N, "input length must be L0 + 2*L1" |
|
|
| |
| bias = torch.full((B, 1, N, N), 0) |
|
|
| rows = torch.arange(L0 + L1, L0 + 2 * L1) |
| rows_token = torch.arange(L0, L0 + L1) |
|
|
| |
| for bi in range((L1 + block_size - 1) // block_size): |
| |
| left_end = L0 + min((bi) * block_size, L1) |
| right_start= L0 + L1 + (left_end - L0) |
|
|
| i_start = bi * block_size |
| i_end = min((bi + 1) * block_size, L1) |
|
|
| block_rows = rows[i_start:i_end] |
| bias[:, :, block_rows.unsqueeze(-1), 0:left_end] = 1 |
| bias[:, :, block_rows.unsqueeze(-1), right_start:(right_start + block_size)] = 1 |
|
|
| block_rows = rows_token[i_start:i_end] |
| left_end = L0 + min((bi + 1) * block_size, L1) |
| bias[:, :, block_rows.unsqueeze(-1), 0:left_end] = 1 |
| |
| if L0 > 0: |
| num_blocks_pre = (L0 + block_size - 1) // block_size |
| for bi in range(num_blocks_pre): |
| |
| row_end = max(L0 - bi * block_size, 0) |
| row_start = max(L0 - (bi + 1) * block_size, 0) |
| if row_end > row_start: |
| block_rows = torch.arange(row_start, row_end) |
| bias[:, :, block_rows.unsqueeze(-1), 0:row_end] = 1 |
| |
| return bias |
|
|
| def process_pad(attn, input_ids, L0, L1, start_pos, pad_id): |
| N = L0 + 2 * L1 |
| device = input_ids.device |
|
|
| cols = torch.arange(N, device=device) |
| key_mask = (cols < start_pos).unsqueeze(0) & (input_ids == pad_id) |
|
|
| |
| attn.masked_fill_(key_mask[:, None, None, :], 0) |
|
|
| |
| A = attn[:, 0] |
| bad = (A.sum(dim=-1) == 0) & (torch.arange(A.size(1), device=A.device).unsqueeze(0) < start_pos) |
| b, r = bad.nonzero(as_tuple=True) |
| A[b, r, :] = 0; A[b, r, r] = 1 |
|
|
| return attn |
|
|
| def one_round_vectorized(input_ids_b, step_map_b, L0, L1, block_size, mask_id): |
| """ |
| Perform a single "round" on one sample b: |
| - For each block, take the minimum non -1 value in step_map. |
| - Create pmask (positions equal to the block minimum). |
| - Create a noise mask for the extended segment (positions >= block minimum). |
| - Mark the chosen minimum positions in step_map as -1 for the next round. |
| |
| Returns: |
| extended_input_ids_b : Tensor with duplicated + masked response segment |
| pmask_b : Boolean mask for tokens selected in this round |
| new_step_map_b : Updated step_map (selected positions set to -1) |
| has_any : Whether any position was selected in this round |
| """ |
| device = input_ids_b.device |
| NB = (L1 + block_size - 1) // block_size |
| pad_len = NB * block_size - L1 |
|
|
| |
| step_pad = torch.full((NB * block_size,), -1, dtype=torch.long, device=device) |
| step_pad[:L1] = step_map_b |
| step_blk = step_pad.view(NB, block_size) |
|
|
| valid = step_blk.ge(0) |
| big = torch.iinfo(step_blk.dtype).max |
| tmp = step_blk.masked_fill(~valid, big) |
| min_vals, _ = tmp.min(dim=1, keepdim=True) |
|
|
| |
| pmask_blk = step_blk.eq(min_vals) & valid |
| if not pmask_blk.any(): |
| |
| return None, None, step_map_b, False |
|
|
| |
| ge_mask_blk = step_blk.ge(min_vals) & valid |
|
|
| |
| pmask_tail = pmask_blk.view(-1)[:L1] |
| ge_mask_tail = ge_mask_blk.view(-1)[:L1] |
|
|
| |
| pmask_b = torch.zeros(L0 + L1, dtype=torch.bool, device=device) |
| pmask_b[L0:] = pmask_tail |
|
|
| |
| tail = input_ids_b[L0:L0+L1].clone() |
| tail[ge_mask_tail] = mask_id |
|
|
| extended_input_ids_b = torch.empty(L0 + L1 + L1, dtype=input_ids_b.dtype, device=device) |
| extended_input_ids_b[:L0+L1] = input_ids_b |
| extended_input_ids_b[L0+L1:] = tail |
|
|
| |
| new_step_map_b = step_map_b.clone() |
| new_step_map_b[pmask_tail] = -1 |
|
|
| return extended_input_ids_b, pmask_b, new_step_map_b, True |
|
|
|
|
| def collapse_k_unique(lst, k: int): |
| if k <= 0: |
| raise ValueError("k must be > 0") |
| uniq = sorted(set(lst)) |
|
|
| mapping = {} |
| n = len(uniq) |
| for idx, val in enumerate(uniq): |
| group = idx // k |
| end_idx = min((group + 1) * k - 1, n - 1) |
| rep = uniq[end_idx] |
| mapping[val] = rep |
| return [mapping[x] for x in lst] |
|
|
| def collect_training_data(config, input_ids, start_pos, pad_id, mask_id, vocab_size=None, post_num=None, step_map_list=None): |
| B, L = input_ids.shape |
| L0 = start_pos |
| L1 = L - L0 |
|
|
| |
|
|
| |
| |
|
|
| if config.training.method == "semi-ar": |
| |
| |
| mask_ratios = config.training.get("mask_ratios", None) |
| if mask_ratios is None: |
| mask_ratios = [1.0] |
| elif isinstance(mask_ratios, (int, float)): |
| mask_ratios = [mask_ratios] |
| else: |
| |
| try: |
| mask_ratios = list(mask_ratios) |
| except (TypeError, AttributeError): |
| pass |
| |
| |
| random_ratio = config.model.get("random_ratio", 0.0) |
| |
| |
| |
| mask_strategy = config.training.get("mask_strategy", "trace") |
| if mask_strategy not in ["trace", "random"]: |
| raise ValueError(f"mask_strategy must be 'trace' or 'random', got '{mask_strategy}'") |
| |
| |
| block_size = config.training.block_size |
| |
| device = input_ids.device |
| |
| |
| |
| |
| mask_ratio_exponent = config.training.get("mask_ratio_exponent", 4.0) |
| mask_ratios_tensor = torch.tensor(mask_ratios, dtype=torch.float32, device=device) |
| weights = mask_ratios_tensor ** mask_ratio_exponent |
| |
| probs = weights / weights.sum() |
| |
| |
| |
| selected_mask_ratios_list = [] |
| |
| |
| |
| sampled_indices = torch.multinomial(probs.unsqueeze(0).expand(B, -1), num_samples=1, replacement=True).squeeze(-1) |
| selected_mask_ratios_list = [mask_ratios[idx.item()] for idx in sampled_indices] |
| |
| |
| |
| step_map_expanded = None |
| if step_map_list is not None and len(step_map_list) > 0: |
| step_map_expanded = [] |
| for b in range(B): |
| sm = step_map_list[b] |
| if isinstance(sm, (list, tuple)): |
| step_map_expanded.append(torch.tensor(sm, dtype=torch.long)) |
| elif isinstance(sm, torch.Tensor): |
| step_map_expanded.append(sm.clone()) |
| else: |
| step_map_expanded.append(torch.tensor(sm, dtype=torch.long)) |
| |
| |
| input_ids_expanded = input_ids |
| expanded_B = B |
| selected_mask_ratios = torch.tensor(selected_mask_ratios_list, device=device, dtype=torch.float32) |
| |
| |
| noise_tail = input_ids_expanded[:, L0:].clone() |
| |
| |
| response_mask = torch.zeros(expanded_B, L1, dtype=torch.bool, device=device) |
| |
| |
| use_trace_masking = (mask_strategy == "trace") and (step_map_expanded is not None) |
| |
| if use_trace_masking: |
| |
| |
| step_map_tensors = [] |
| for sm in step_map_expanded: |
| if isinstance(sm, (list, tuple)): |
| step_map_tensors.append(torch.tensor(sm, dtype=torch.long)) |
| elif isinstance(sm, torch.Tensor): |
| step_map_tensors.append(sm) |
| else: |
| step_map_tensors.append(torch.tensor(sm, dtype=torch.long)) |
| |
| |
| |
| max_len = max(sm.shape[0] for sm in step_map_tensors) |
| step_map_padded = [] |
| for sm in step_map_tensors: |
| sm_len = sm.shape[0] |
| if sm_len < max_len: |
| |
| padding = torch.full((max_len - sm_len,), 999999, dtype=sm.dtype) |
| sm = torch.cat([sm, padding], dim=0) |
| elif sm_len > max_len: |
| sm = sm[:max_len] |
| step_map_padded.append(sm) |
| |
| step_map = torch.stack(step_map_padded, dim=0).to(device) |
| |
| |
| if step_map.shape[1] > L1: |
| step_map = step_map[:, :L1] |
| elif step_map.shape[1] < L1: |
| |
| pad_len = L1 - step_map.shape[1] |
| padding = torch.full((expanded_B, pad_len), 999999, dtype=step_map.dtype, device=device) |
| step_map = torch.cat([step_map, padding], dim=1) |
| |
| |
| NB = (L1 + block_size - 1) // block_size |
| for b in range(expanded_B): |
| mask_ratio = selected_mask_ratios[b].item() |
| step_map_b = step_map[b] |
| |
| |
| for bi in range(NB): |
| block_start = bi * block_size |
| block_end = min((bi + 1) * block_size, L1) |
| block_len = block_end - block_start |
| |
| |
| block_step_map = step_map_b[block_start:block_end] |
| block_indices = torch.arange(block_start, block_end, device=device) |
| |
| |
| valid_mask = block_step_map < 999999 |
| valid_block_indices = block_indices[valid_mask] |
| valid_block_step_map = block_step_map[valid_mask] |
| |
| if len(valid_block_indices) > 0: |
| |
| sorted_order = torch.argsort(valid_block_step_map) |
| sorted_valid_indices = valid_block_indices[sorted_order] |
| |
| |
| num_to_mask_in_block = int(len(sorted_valid_indices) * mask_ratio) |
| if num_to_mask_in_block > 0: |
| |
| mask_indices_in_block = sorted_valid_indices[:num_to_mask_in_block] |
| response_mask[b, mask_indices_in_block] = True |
| else: |
| |
| NB = (L1 + block_size - 1) // block_size |
| for b in range(expanded_B): |
| mask_ratio = selected_mask_ratios[b].item() |
| |
| |
| for bi in range(NB): |
| block_start = bi * block_size |
| block_end = min((bi + 1) * block_size, L1) |
| block_len = block_end - block_start |
| |
| |
| num_to_mask_in_block = int(block_len * mask_ratio) |
| if num_to_mask_in_block > 0: |
| |
| block_positions = torch.randperm(block_len, device=device)[:num_to_mask_in_block] |
| mask_indices_in_block = block_start + block_positions |
| response_mask[b, mask_indices_in_block] = True |
| |
| |
| p_mask = torch.cat([ |
| torch.zeros(expanded_B, L0, dtype=torch.bool, device=device), |
| response_mask |
| ], dim=1) |
| |
| |
| if random_ratio > 0 and vocab_size is not None: |
| |
| |
| masked_positions = response_mask |
| |
| |
| random_mask = torch.zeros(expanded_B, L1, dtype=torch.bool, device=device) |
| |
| |
| for b in range(expanded_B): |
| masked_idx = torch.where(masked_positions[b])[0] |
| if len(masked_idx) > 0: |
| num_random = max(1, int(len(masked_idx) * random_ratio)) |
| num_random = min(num_random, len(masked_idx)) |
| if num_random > 0: |
| |
| random_idx = masked_idx[torch.randperm(len(masked_idx), device=device)[:num_random]] |
| random_mask[b, random_idx] = True |
| |
| mask_token_mask = masked_positions & (~random_mask) |
| |
| |
| if random_mask.any(): |
| num_random = random_mask.sum().item() |
| random_tokens = torch.randint(0, vocab_size, (num_random,), |
| device=device, dtype=noise_tail.dtype) |
| noise_tail[random_mask] = random_tokens |
| |
| |
| if mask_token_mask.any(): |
| noise_tail[mask_token_mask] = mask_id |
| else: |
| |
| noise_tail[response_mask] = mask_id |
| |
| |
| extended_input_ids = torch.cat([input_ids_expanded, noise_tail], dim=1) |
|
|
| else: |
| raise ValueError(f"Method {config.training.method} not supported") |
| |
| pad_resp = (extended_input_ids[:, :L] == pad_id) & p_mask |
| if post_num is not None: |
| cum_pad = torch.cumsum(pad_resp.int(), dim=1) |
| p_mask &= ~(pad_resp & (cum_pad > post_num)) |
| |
| labels = extended_input_ids[:, :L].clone() |
|
|
| idx = torch.arange(L).unsqueeze(0).expand(extended_input_ids.shape[0], -1) |
| valid = (idx >= start_pos) | extended_input_ids[:, :L].ne(pad_id) |
| tok_idx = valid.long().cumsum(dim=-1) - 1 |
| tok_idx = tok_idx.masked_fill(~valid, 1) |
| tok_idx_resp = tok_idx[:, start_pos:] |
| tok_idx_ext = torch.cat([tok_idx, tok_idx_resp], dim=1) |
|
|
| keep = p_mask.view(p_mask.size(0), -1).any(dim=1) |
|
|
| extended_input_ids = extended_input_ids[keep] |
| p_mask = p_mask[keep] |
| tok_idx_ext = tok_idx_ext[keep] |
| labels = labels[keep] |
|
|
| return extended_input_ids, p_mask, tok_idx_ext, labels |