feat: add paired DeepSeek routing length control

This commit is contained in:
wuyang
2026-07-29 16:17:18 +08:00
parent 059458f8e3
commit 9ca08501a8
15 changed files with 318874 additions and 69 deletions
+53 -11
View File
@@ -69,6 +69,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--layers", type=int, default=7)
parser.add_argument("--bootstrap", type=int, default=2000)
parser.add_argument("--seed", type=int, default=20260729)
parser.add_argument("--sample-salt", default=SAMPLE_SALT)
parser.add_argument("--eligibility-min-tokens", type=int, default=None)
parser.add_argument("--device", default="cuda")
parser.add_argument("--captured-at", default=None)
return parser.parse_args()
@@ -175,6 +177,8 @@ def select_corpus(
candidates: dict[str, list[dict[str, str]]],
per_domain: int,
max_tokens: int,
sample_salt: str,
eligibility_min_tokens: int | None,
) -> tuple[list[dict[str, Any]], dict[str, dict[str, int]]]:
selected: list[dict[str, Any]] = []
counts: dict[str, dict[str, int]] = {}
@@ -183,21 +187,31 @@ def select_corpus(
texts = [row["text"] for row in rows]
encoded: list[list[int]] = []
for start in range(0, len(texts), 512):
result = tokenizer(
texts[start : start + 512],
add_special_tokens=True,
truncation=True,
max_length=max_tokens,
padding=False,
)
tokenizer_kwargs: dict[str, Any] = {
"add_special_tokens": True,
"padding": False,
}
if eligibility_min_tokens is None:
tokenizer_kwargs.update(
{
"truncation": True,
"max_length": max_tokens,
}
)
else:
tokenizer_kwargs["truncation"] = False
result = tokenizer(texts[start : start + 512], **tokenizer_kwargs)
encoded.extend(result.input_ids)
eligible = []
for row, token_ids in zip(rows, encoded, strict=True):
if len(token_ids) < MIN_TOKENS[domain]:
minimum = eligibility_min_tokens or MIN_TOKENS[domain]
if len(token_ids) < minimum:
continue
source_tokens = len(token_ids)
token_ids = token_ids[:max_tokens]
rank = hashlib.sha256(
f"{SAMPLE_SALT}|{domain}|{row['id']}".encode()
f"{sample_salt}|{domain}|{row['id']}".encode()
).hexdigest()
eligible.append(
{
@@ -209,6 +223,11 @@ def select_corpus(
"characters": len(row["text"]),
"token_ids": token_ids,
"tokens": len(token_ids),
**(
{"source_tokens": source_tokens}
if eligibility_min_tokens is not None
else {}
),
"selection_rank": rank,
}
)
@@ -509,6 +528,10 @@ def main() -> None:
raise ValueError("per-domain sample must be at least two")
if args.bootstrap < 100:
raise ValueError("bootstrap replicates must be at least 100")
if args.eligibility_min_tokens is not None and args.eligibility_min_tokens < 2:
raise ValueError("eligibility minimum must be at least two tokens")
if not args.sample_salt.strip():
raise ValueError("sample salt must not be empty")
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
@@ -532,6 +555,8 @@ def main() -> None:
candidates,
args.per_domain,
args.max_tokens,
args.sample_salt,
args.eligibility_min_tokens,
)
batches = make_batches(samples, args.batch_size, tokenizer.pad_token_id)
device = torch.device(args.device)
@@ -658,6 +683,11 @@ def main() -> None:
"text_sha256": sample["text_sha256"],
"characters": sample["characters"],
"tokens": sample["tokens"],
**(
{"source_tokens": sample["source_tokens"]}
if "source_tokens" in sample
else {}
),
}
for sample in samples
]
@@ -754,11 +784,23 @@ def main() -> None:
"corpus_contract": {
"domains": list(DOMAIN_ORDER),
"domain_labels": DOMAIN_LABELS,
"sample_salt": SAMPLE_SALT,
"sample_salt": args.sample_salt,
"selection": "ascending SHA256(salt|domain|source_id), then source_id",
"per_domain": args.per_domain,
"max_tokens": args.max_tokens,
"minimum_tokens": MIN_TOKENS,
"minimum_tokens": (
{domain: args.eligibility_min_tokens for domain in DOMAIN_ORDER}
if args.eligibility_min_tokens is not None
else MIN_TOKENS
),
**(
{
"uniform_eligibility_min_tokens": args.eligibility_min_tokens,
"matched_length_control": True,
}
if args.eligibility_min_tokens is not None
else {}
),
"special_tokens": True,
"truncation": "right",
"counts": corpus_counts,