File size: 2,194 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
"""
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()