Download sdg/generate.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 2.19 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/generate.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/generate.py
-
curl -L -o generate.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/generate.py
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() | |