feat: add paired DeepSeek routing length control
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user