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

219 lines
8.8 KiB
Python

# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
"""
Data curation pipeline: generate responses from a dataset using vLLM.
Each worker (identified by --rank) processes a disjoint shard of the input
dataset, generates responses via vLLM offline inference, and writes results
as Arrow IPC files (one per batch) into a rank-specific output directory.
Checkpointing allows resuming from the last completed batch.
Standalone:
python data_curation/pipeline.py \
--model Qwen/Qwen3-4B \
--input data.jsonl \
--output-dir output/
Multi-GPU (one model per GPU):
See run_curation.sh for the recommended launch pattern.
"""
import argparse
import json
import os
import pickle
from pathlib import Path
import pandas as pd
import pyarrow as pa
import pyarrow.ipc as ipc
from tqdm import tqdm
from vllm import LLM, SamplingParams
# ---------------------------------------------------------------------------
# Data I/O
# ---------------------------------------------------------------------------
def load_dataset(path: str) -> list[dict]:
"""Load a .jsonl or .parquet dataset into a list of dicts."""
if path.endswith(".parquet"):
df = pd.read_parquet(path)
records = df.to_dict("records")
for record in records:
if "prompt" in record and hasattr(record["prompt"], "tolist"):
record["prompt"] = record["prompt"].tolist()
return records
elif path.endswith(".jsonl"):
with open(path) as f:
return [json.loads(line) for line in f]
else:
raise ValueError(f"Unsupported format: {path}. Use .jsonl or .parquet.")
def save_batch_arrow(rows: list[dict], path: str) -> None:
"""Write a list of dicts as an Arrow IPC file."""
table = pa.Table.from_pandas(pd.DataFrame(rows))
with pa.OSFile(path, "wb") as sink:
with ipc.new_file(sink, table.schema) as writer:
writer.write_table(table)
# ---------------------------------------------------------------------------
# Core pipeline
# ---------------------------------------------------------------------------
def run_curation(args: argparse.Namespace) -> None:
tag = f"[Rank {args.rank}/{args.world_size}]"
# ── Load & shard dataset ──────────────────────────────────────────────
print(f"{tag} Loading dataset: {args.input}")
dataset = load_dataset(args.input)
if args.num_samples is not None:
dataset = dataset[: args.num_samples]
print(f"{tag} Debug mode: limiting to {args.num_samples} samples")
if args.world_size > 1:
dataset = dataset[args.rank :: args.world_size]
print(f"{tag} Assigned {len(dataset)} samples")
# ── Output directory ──────────────────────────────────────────────────
if args.world_size > 1:
output_dir = Path(args.output_dir) / f"rank{args.rank:05d}"
else:
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
# ── Checkpoint ────────────────────────────────────────────────────────
ckpt_dir = Path(args.checkpoint_dir)
ckpt_dir.mkdir(parents=True, exist_ok=True)
ckpt_file = ckpt_dir / f"rank{args.rank:05d}.pkl"
start_idx = 0
if ckpt_file.exists():
with open(ckpt_file, "rb") as f:
start_idx = pickle.load(f)["next_idx"]
print(f"{tag} Resuming from index {start_idx}")
# ── Model ─────────────────────────────────────────────────────────────
print(f"{tag} Loading model: {args.model} (tp={args.tensor_parallel_size})")
llm = LLM(
model=args.model,
tensor_parallel_size=args.tensor_parallel_size,
trust_remote_code=True,
)
sampling_params = SamplingParams(
temperature=args.temperature,
top_p=args.top_p,
max_tokens=args.max_tokens,
n=args.num_responses,
)
# ── Batch loop ────────────────────────────────────────────────────────
total_batches = (len(dataset) + args.batch_size - 1) // args.batch_size
total_saved = 0
print(f"{tag} Processing {len(dataset)} prompts, batch_size={args.batch_size}, "
f"total_batches={total_batches}")
for batch_start in range(start_idx, len(dataset), args.batch_size):
batch_end = min(batch_start + args.batch_size, len(dataset))
batch = dataset[batch_start:batch_end]
batch_idx = batch_start // args.batch_size
prompts = [item["prompt"] for item in batch]
print(f"{tag} Batch {batch_idx + 1}/{total_batches} "
f"({batch_end - batch_start} samples) ...")
outputs = llm.chat(prompts, sampling_params)
# Build results
rows = []
for item, output in zip(batch, outputs):
for completion in output.outputs:
text = completion.text
# Ensure <think> tag is present
if "</think>" in text and not text.strip().startswith("<think>"):
text = "<think>\n" + text
messages = item["prompt"] + [{"role": "assistant", "content": text}]
rows.append({
"messages": messages,
"tokens": len(completion.token_ids),
})
# Save Arrow file
arrow_path = output_dir / f"data-{batch_idx:05d}-of-{total_batches:05d}.arrow"
save_batch_arrow(rows, str(arrow_path))
total_saved += len(rows)
# Save checkpoint
with open(ckpt_file, "wb") as f:
pickle.dump({"next_idx": batch_end}, f)
print(f"{tag} Saved {arrow_path.name} (total: {total_saved})")
# ── Cleanup ───────────────────────────────────────────────────────────
if ckpt_file.exists():
ckpt_file.unlink()
print(f"{tag} Done! {total_saved} samples → {output_dir}/")
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Generate responses from a dataset using vLLM offline inference.",
)
# Required
p.add_argument("--model", type=str, required=True,
help="HuggingFace model name or path.")
p.add_argument("--input", type=str, required=True,
help="Input dataset (.jsonl or .parquet).")
p.add_argument("--output-dir", type=str, required=True,
help="Root output directory. Each rank writes to a subdirectory.")
# Generation
p.add_argument("--max-tokens", type=int, default=16384,
help="Max new tokens per response (default: 16384).")
p.add_argument("--temperature", type=float, default=0.7,
help="Sampling temperature (default: 0.7).")
p.add_argument("--top-p", type=float, default=0.9,
help="Nucleus sampling top-p (default: 0.9).")
p.add_argument("--num-responses", type=int, default=1,
help="Number of responses per prompt (default: 1).")
p.add_argument("--batch-size", type=int, default=32,
help="Prompts per vLLM batch call (default: 32).")
# Parallelism
p.add_argument("--tensor-parallel-size", type=int, default=1,
help="vLLM tensor-parallel size (default: 1).")
p.add_argument("--rank", type=int, default=None,
help="Worker rank (auto-detected from env if omitted).")
p.add_argument("--world-size", type=int, default=None,
help="Total workers (auto-detected from env if omitted).")
# Misc
p.add_argument("--num-samples", type=int, default=None,
help="Limit total samples before sharding (for debugging).")
p.add_argument("--checkpoint-dir", type=str, default="checkpoints",
help="Directory for per-rank checkpoint files (default: checkpoints).")
args = p.parse_args()
# Auto-detect rank / world_size from environment (torchrun, etc.)
if args.rank is None:
args.rank = int(os.environ.get("RANK", os.environ.get("LOCAL_RANK", 0)))
if args.world_size is None:
args.world_size = int(os.environ.get("WORLD_SIZE", 1))
return args
if __name__ == "__main__":
run_curation(parse_args())