Files
myLightningOPD/data_curation/prepare_lightning_opd.py
ModelHub XC d4e0a1af66 初始化项目,由ModelHub XC社区提供模型
Model: ayh015/myLightningOPD
Source: Original Platform
2026-08-27 23:50:14 +08:00

248 lines
9.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""
Prepare Lightning OPD parquet from student rollout data.
Phase 1 tokenize (CPU-friendly):
Reads student rollout parquet, builds prompt via chat template,
tokenizes responses, truncates to --max-response-len, writes intermediate
parquet WITHOUT teacher logprobs.
Phase 2 precompute teacher logprobs (requires GPU / teacher sglang server):
Reads the intermediate parquet produced in Phase 1, sends each
(prompt + response) sequence to the teacher sglang server, stores
per-token response logprobs back into the metadata, writes the final
parquet.
Usage (Phase 1, CPU node):
python3 data_curation/prepare_lightning_opd.py \\
--tokenizer-path checkpoints/sft \\
--input-parquet data/rollouts/rollouts.parquet \\
--output-dir data/lightning_opd
Usage (Phase 2, GPU node with teacher sglang running):
python3 data_curation/prepare_lightning_opd.py \\
--tokenizer-path checkpoints/sft \\
--input-parquet data/rollouts/rollouts.parquet \\
--output-dir data/lightning_opd \\
--compute-teacher-logprobs \\
--teacher-url http://127.0.0.1:13141/generate
"""
import argparse
import asyncio
from pathlib import Path
import aiohttp
import pandas as pd
from transformers import AutoTokenizer
from tqdm import tqdm
def parse_args():
parser = argparse.ArgumentParser(
description="Prepare Lightning OPD parquet data (tokenize + optional teacher logprobs)."
)
parser.add_argument(
"--tokenizer-path", type=str, required=True,
help="Path to HuggingFace tokenizer (e.g. the student SFT checkpoint).",
)
parser.add_argument(
"--input-parquet", type=str, required=True,
help="Path to student rollout parquet. Expected columns: messages (list[dict]), tokens (int).",
)
parser.add_argument(
"--output-dir", type=str, required=True,
help="Directory where intermediate and final parquet files are written.",
)
parser.add_argument(
"--max-response-len", type=int, default=4096,
help="Maximum response token length; longer responses are truncated (default: 4096).",
)
parser.add_argument(
"--compute-teacher-logprobs", action="store_true",
help="Run Phase 2: compute teacher logprobs via a running sglang server.",
)
parser.add_argument(
"--teacher-url", type=str, default="http://127.0.0.1:13141/generate",
help="Teacher sglang server URL (default: http://127.0.0.1:13141/generate).",
)
parser.add_argument(
"--concurrency", type=int, default=64,
help="Number of concurrent requests to teacher sglang server (default: 64).",
)
return parser.parse_args()
# ── Phase 1: tokenize ────────────────────────────────────────────────────────
def phase1_tokenize(args, intermediate_path: Path):
print(f"[Phase 1] Loading tokenizer from {args.tokenizer_path}")
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_path, trust_remote_code=True)
print(f"[Phase 1] Loading input parquet: {args.input_parquet}")
df = pd.read_parquet(args.input_parquet)
print(f"[Phase 1] Total rows: {len(df)}")
rows_out = []
truncated = 0
skipped = 0
for row in tqdm(df.itertuples(), total=len(df), desc="Tokenizing"):
messages = row.messages
user_messages = [m for m in messages if m["role"] != "assistant"]
prompt_str = tokenizer.apply_chat_template(
user_messages, tokenize=False, add_generation_prompt=True, enable_thinking=True
)
assistant_msg = None
for msg in messages:
if msg["role"] == "assistant":
assistant_msg = msg["content"]
break
if assistant_msg is None:
skipped += 1
continue
response_ids = tokenizer.encode(assistant_msg, add_special_tokens=False)
if len(response_ids) > args.max_response_len:
truncated += 1
response_ids = response_ids[:args.max_response_len]
assistant_msg = tokenizer.decode(response_ids, skip_special_tokens=False)
rows_out.append({
"prompt": prompt_str,
"label": "0",
"metadata": {
"is_lightning_opd": True,
"response_tokens": response_ids,
"loss_mask": [1] * len(response_ids),
"response": assistant_msg,
},
})
print(f"[Phase 1] Rows written: {len(rows_out)}, "
f"truncated to {args.max_response_len}: {truncated}, skipped: {skipped}")
df_out = pd.DataFrame(rows_out)
intermediate_path.parent.mkdir(parents=True, exist_ok=True)
df_out.to_parquet(intermediate_path, index=False)
print(f"[Phase 1] Saved to {intermediate_path}")
# ── Phase 2: precompute teacher logprobs ─────────────────────────────────────
async def _fetch_logprobs(
session: aiohttp.ClientSession,
teacher_url: str,
full_ids: list[int],
response_len: int,
) -> list[float]:
"""Call teacher sglang server and return per-token logprobs for the response portion."""
payload = {
"input_ids": full_ids,
"sampling_params": {
"temperature": 0,
"max_new_tokens": 0,
"skip_special_tokens": False,
},
"return_logprob": True,
"logprob_start_len": 0,
}
async with session.post(teacher_url, json=payload) as resp:
resp.raise_for_status()
ret = await resp.json()
all_lps = ret["meta_info"]["input_token_logprobs"]
response_lps = [float(item[0]) for item in all_lps[1:]][-response_len:]
assert len(response_lps) == response_len, (
f"Expected {response_len} logprobs, got {len(response_lps)}"
)
return response_lps
async def _process_all(args, tokenizer, rows: list[dict]) -> list[list[float]]:
"""Process all rows concurrently with a live progress bar, preserving order."""
semaphore = asyncio.Semaphore(args.concurrency)
connector = aiohttp.TCPConnector(limit=args.concurrency)
results = [None] * len(rows)
async def bounded_fetch(idx: int, full_ids: list[int], response_len: int):
async with semaphore:
result = await _fetch_logprobs(session, args.teacher_url, full_ids, response_len)
results[idx] = result
pbar.update(1)
async with aiohttp.ClientSession(connector=connector) as session:
with tqdm(total=len(rows), desc="[Phase 2] Teacher logprobs") as pbar:
tasks = []
for idx, row in enumerate(rows):
meta = row["metadata"]
prompt_ids = tokenizer.encode(row["prompt"], add_special_tokens=False)
response_ids = [int(x) for x in meta["response_tokens"]]
full_ids = prompt_ids + response_ids
tasks.append(bounded_fetch(idx, full_ids, len(response_ids)))
await asyncio.gather(*tasks)
return results
def phase2_logprobs(args, intermediate_path: Path, output_path: Path):
print(f"[Phase 2] Loading intermediate parquet: {intermediate_path}")
df = pd.read_parquet(intermediate_path)
rows = df.to_dict(orient="records")
print(f"[Phase 2] Total rows: {len(rows)}")
print(f"[Phase 2] Loading tokenizer from {args.tokenizer_path}")
tokenizer = AutoTokenizer.from_pretrained(args.tokenizer_path, trust_remote_code=True)
print(f"[Phase 2] Computing teacher logprobs via {args.teacher_url} "
f"(concurrency={args.concurrency})")
all_logprobs = asyncio.run(_process_all(args, tokenizer, rows))
for row, lps in zip(rows, all_logprobs):
row["metadata"]["teacher_log_probs"] = lps
df_out = pd.DataFrame(rows)
output_path.parent.mkdir(parents=True, exist_ok=True)
df_out.to_parquet(output_path, index=False)
print(f"[Phase 2] Saved to {output_path}")
# Sanity check
df_check = pd.read_parquet(output_path)
row0 = df_check.iloc[0]
meta = row0["metadata"]
print("\n[Phase 2] Sanity check row 0:")
print(f" prompt[:80]: {row0['prompt'][:80]}")
print(f" label: {row0['label']}")
print(f" len(response_tokens): {len(meta['response_tokens'])}")
print(f" len(teacher_log_probs): {len(meta['teacher_log_probs'])}")
print(f" teacher_log_probs[:5]: {meta['teacher_log_probs'][:5]}")
# ── Entry point ───────────────────────────────────────────────────────────────
def main():
args = parse_args()
output_dir = Path(args.output_dir)
input_stem = Path(args.input_parquet).stem
intermediate_path = output_dir / f"{input_stem}-lightning-opd.parquet"
output_path = output_dir / f"{input_stem}-lightning-opd-precomputed.parquet"
if args.compute_teacher_logprobs:
if not intermediate_path.exists():
print("[INFO] Intermediate parquet not found, running Phase 1 first.")
phase1_tokenize(args, intermediate_path)
phase2_logprobs(args, intermediate_path, output_path)
else:
phase1_tokenize(args, intermediate_path)
print(f"\n[INFO] To add teacher logprobs, re-run with --compute-teacher-logprobs "
f"after starting the teacher sglang server.")
if __name__ == "__main__":
main()