svd-code / sdg /generate.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
2.19 kB
"""
CLI entry point for the SDG pipeline.
Usage:
python -m sdg.generate --config sdg/configs/qwen3_4b_mot.yaml
python -m sdg.generate --config sdg/configs/qwen3_4b_openthoughts4.yaml --limit 100
"""
import argparse
from sdg.config import SDGConfig
from sdg.pipeline import run_pipeline
def main() -> None:
parser = argparse.ArgumentParser(
description="Standalone SDG pipeline (generate + UQ validation)"
)
parser.add_argument(
"--config", required=True, help="Path to YAML config file"
)
# CLI overrides
parser.add_argument("--model", dest="model_name")
parser.add_argument("--dataset", dest="dataset_hf_id")
parser.add_argument("--dataset-config", dest="dataset_hf_config")
parser.add_argument("--limit", type=int)
parser.add_argument("--num-generations", type=int)
parser.add_argument("--num-validation-votes", type=int)
parser.add_argument("--max-rounds", type=int)
parser.add_argument("--output-dir")
parser.add_argument("--tensor-parallel-size", type=int)
# Sharding
parser.add_argument("--shard-id", type=int, default=None, help="Shard index (0-based)")
parser.add_argument("--num-shards", type=int, default=None, help="Total number of shards")
args = parser.parse_args()
# Load base config from YAML
config = SDGConfig.from_yaml(args.config)
# Apply CLI overrides
overrides = {
"model_name": args.model_name,
"dataset_hf_id": args.dataset_hf_id,
"dataset_hf_config": args.dataset_hf_config,
"limit": args.limit,
"num_generations": args.num_generations,
"num_validation_votes": args.num_validation_votes,
"max_rounds": args.max_rounds,
"output_dir": args.output_dir,
"tensor_parallel_size": args.tensor_parallel_size,
}
for key, value in overrides.items():
if value is not None:
setattr(config, key, value)
if (args.shard_id is None) != (args.num_shards is None):
parser.error("--shard-id and --num-shards must be provided together")
run_pipeline(config, shard_id=args.shard_id, num_shards=args.num_shards)
if __name__ == "__main__":
main()