Spaces:
Running on Zero
Running on Zero
File size: 10,086 Bytes
a0e2620 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | """Online corruption; validation is deterministic by seed and example index."""
from __future__ import annotations
import torch
def _legacy_structural_noise(tokens: torch.Tensor, mask_token_id: int, generator: torch.Generator | None = None) -> torch.Tensor:
"""Apply the historical LAD answer corruption to one unpadded answer."""
corrupted = tokens.clone()
length = len(corrupted)
if not length:
return corrupted
# One in ten examples starts from an entirely masked answer.
if torch.rand((), generator=generator).item() < 0.1:
return torch.full_like(corrupted, mask_token_id)
noise_prob = torch.rand((), generator=generator).item()
# The original routine applies this independent masking block half the time.
if torch.rand((), generator=generator).item() < 0.5:
mask_fraction = torch.rand((), generator=generator).item() * 0.5
count = int(length * mask_fraction)
if count:
indices = torch.randperm(length, generator=generator)[:count]
corrupted[indices] = mask_token_id
if length <= 2:
return corrupted
swap_mask = torch.rand(length - 1, generator=generator) < (noise_prob / 4)
for index in torch.where(swap_mask)[0].tolist():
value = corrupted[index].clone()
corrupted[index] = corrupted[index + 1]
corrupted[index + 1] = value
duplicate_mask = torch.rand(length, generator=generator) < (noise_prob / 4)
duplicate_indices = torch.where(duplicate_mask)[0]
if len(duplicate_indices):
directions = torch.randint(0, 2, (len(duplicate_indices),), generator=generator)
backward = duplicate_indices[(directions == 0) & (duplicate_indices > 0)]
forward = duplicate_indices[(directions == 1) & (duplicate_indices < length - 1)]
# NumPy advanced indexing copies the RHS for each assignment. The
# legacy code applies backward copies first, then forward copies.
corrupted[backward] = corrupted[backward - 1]
corrupted[forward] = corrupted[forward + 1]
if torch.rand((), generator=generator).item() < (noise_prob / 4):
span_length = int(torch.randint(1, min(3, length) + 1, (), generator=generator).item())
shift = int(torch.randint(1, 5, (), generator=generator).item())
direction = -1 if torch.randint(0, 2, (), generator=generator).item() == 0 else 1
start = int(torch.randint(0, length - span_length + 1, (), generator=generator).item())
span = corrupted[start : start + span_length].clone()
target = max(0, start - shift) if direction == -1 else min(length - span_length, start + shift)
corrupted[target : target + span_length] = span
return corrupted
def apply_corruption(batch, mask_token_id, mode, structured_loss_behavior, eos_padding_loss, t_min, seed, deterministic,
frontier_masking_probability=0.0, frontier_masking_epsilon=0.03, frontier_masking_tau=3.0,
frontier_padding_mode="iid"):
"""Apply configured corruption and choose the positions used for loss."""
answer = batch["answer_mask"] & ~batch["padding_mask"]
eos_padding = batch["padding_mask"]
supervised = answer | eos_padding if eos_padding_loss else answer
if mode == "structured":
online = batch.pop("structured_online", torch.zeros(answer.shape[0], dtype=torch.bool))
for row, needs_noise in enumerate(online.tolist()):
if not needs_noise:
continue
generator = None
if deterministic:
generator = torch.Generator(device="cpu").manual_seed(seed + int(batch["example_index"][row]))
positions = torch.where(answer[row])[0]
batch["input_ids"][row, positions] = _legacy_structural_noise(
batch["labels"][row, positions], mask_token_id, generator
)
if structured_loss_behavior == "all_answer_tokens":
loss_mask = supervised
elif structured_loss_behavior == "corrupted_answer_tokens":
loss_mask = supervised & (batch["input_ids"] != batch["labels"])
elif structured_loss_behavior == "all_tokens":
# EOS padding is controlled separately so it can be compared with
# answer-only objectives without changing the primary loss mode.
loss_mask = ~batch["padding_mask"] | eos_padding if eos_padding_loss else ~batch["padding_mask"]
else:
raise ValueError(f"Unknown structured_loss_behavior={structured_loss_behavior}; expected all_answer_tokens, corrupted_answer_tokens, or all_tokens")
batch["loss_mask"] = loss_mask
batch["sampled_t"] = torch.full((answer.shape[0],), float("nan"))
return batch
batch.pop("structured_online", None)
noised = batch["labels"].clone()
# When selected for loss, EOS padding is a real denoising target:
# it must sometimes be replaced by MASK so the same-position loss teaches
# the model to *produce* EOS rather than merely copy a visible one. Other
# mask-only objectives do not supervise padding and therefore leave it
# untouched.
eligible_mask_only = supervised
selected = torch.zeros_like(eligible_mask_only)
# The IID branch keeps its original inverse-t weighting. Frontier rows
# correct that weighting by t / p(position), leaving raw CE metrics intact.
token_weights = torch.ones_like(noised, dtype=torch.float32) if frontier_masking_probability > 0 else None
ts = []
for row, index in enumerate(batch["example_index"].tolist()):
generator = None
if deterministic:
generator = torch.Generator(device="cpu").manual_seed(seed + index)
t = torch.empty((), dtype=torch.float32).uniform_(t_min, 1.0, generator=generator).item()
eligible = torch.where(eligible_mask_only[row])[0]
if len(eligible):
use_frontier = frontier_masking_probability > 0 and torch.rand((), generator=generator).item() < frontier_masking_probability
if use_frontier:
# The genuine terminating EOS belongs to answer when enabled;
# repeated EOS padding must never move the answer's frontier.
answer_positions = torch.where(answer[row])[0]
positions = torch.arange(len(answer_positions), device=noised.device, dtype=torch.float32)
frontier = len(answer_positions) * (1.0 - t)
probabilities = frontier_masking_epsilon + (1.0 - 2.0 * frontier_masking_epsilon) * torch.sigmoid(
(positions - frontier) / frontier_masking_tau
)
draw = torch.rand(len(answer_positions), generator=generator, device=noised.device) < probabilities
selected[row, answer_positions[draw]] = True
token_weights[row, answer_positions] = t / probabilities
# Draw padding after the answer, so adding padding or changing
# its supervision cannot change this example's answer masks.
if eos_padding_loss:
padding_positions = torch.where(eos_padding[row])[0]
if frontier_padding_mode == "frontier":
# Continue the answer's positional frontier into
# padding without letting padding move the frontier.
padding_offsets = torch.arange(
len(answer_positions),
len(answer_positions) + len(padding_positions),
device=noised.device,
dtype=torch.float32,
)
padding_probabilities = frontier_masking_epsilon + (
1.0 - 2.0 * frontier_masking_epsilon
) * torch.sigmoid((padding_offsets - frontier) / frontier_masking_tau)
else:
# Historical behavior used by the successful original
# run: every EOS-padding position is IID-masked at t.
padding_probabilities = torch.full(
(len(padding_positions),), t, device=noised.device
)
padding_draw = (
torch.rand(len(padding_positions), generator=generator, device=noised.device)
< padding_probabilities
)
selected[row, padding_positions[padding_draw]] = True
if frontier_padding_mode == "frontier":
token_weights[row, padding_positions] = t / padding_probabilities
else:
draw = torch.rand(len(eligible), generator=generator) < t
if not draw.any():
pick = torch.randint(len(eligible), (1,), generator=generator)
draw[pick] = True
selected[row, eligible[draw]] = True
ts.append(t)
noised[selected] = mask_token_id
batch["input_ids"] = noised
if structured_loss_behavior == "all_tokens":
# Keep mask-only's stochastic inputs, but train against every target
# position just like the legacy full-sequence objective. In this mode
# training.py also disables inverse-t weighting.
batch["loss_mask"] = ~batch["padding_mask"] | eos_padding if eos_padding_loss else ~batch["padding_mask"]
elif structured_loss_behavior in {"all_answer_tokens", "corrupted_answer_tokens"}:
# mask_only's historical objective supervises the positions actually
# corrupted in the input; both names retain that behavior here.
batch["loss_mask"] = selected & supervised
else:
raise ValueError(f"Unknown structured_loss_behavior={structured_loss_behavior}; expected all_answer_tokens, corrupted_answer_tokens, or all_tokens")
batch["sampled_t"] = torch.tensor(ts, dtype=torch.float32)
if token_weights is not None:
batch["token_loss_weights"] = token_weights
return batch
|