File size: 3,168 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
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
"""
SDGConfig dataclass and YAML serialization.
"""

from dataclasses import dataclass, field, fields, asdict
from pathlib import Path

import yaml


@dataclass
class SDGConfig:
    # Model
    model_name: str = "Qwen/Qwen3-4B"
    is_instruction_tuned: bool = True
    tensor_parallel_size: int = 1
    max_model_len: int = 32768
    gpu_memory_utilization: float = 0.95

    # Dataset
    dataset_hf_id: str = "open-thoughts/OpenThoughts-114k-math"
    dataset_hf_config: str | None = None  # HF dataset config/subset name
    dataset_hf_revision: str | None = None  # HF dataset revision (commit hash)
    dataset_split: str = "train"
    dataset_format: str = "openthoughts4"  # "openthoughts4" | "messages"
    check_boxed: bool = True
    limit: int | None = None

    # Generation
    num_generations: int = 4
    gen_temperature: float = 0.8
    gen_top_p: float = 0.95
    gen_max_tokens: int = 32768
    gen_batch_size: int = 64

    # Validation
    validation_type: str | None = "uq"  # "uq" | "simple" | null (None = skip validation)
    selection_policy: str = "first_valid"  # "first_valid" | "all_valid"
    num_validation_votes: int = 3
    val_temperature: float = 0.6
    val_top_p: float = 0.95
    val_max_tokens: int = 32768
    val_batch_size: int = 64

    # Output
    experiment_name: str = ""
    output_dir: str = "output"

    # MongoDB cache
    mongo_uri: str = ""

    # Sharding
    num_shards: int = 1

    # Machines to exclude when launching Slurm jobs
    machines_to_exclude: list[str] = field(default_factory=list)

    # ── derived properties ──────────────────────────────────────────────

    @property
    def run_name(self) -> str:
        if not self.experiment_name:
            raise ValueError("experiment_name must be set in the config YAML")
        return self.experiment_name

    @property
    def mongo_db_name(self) -> str:
        return "sdg_cache"

    @property
    def output_path(self) -> Path:
        return Path(self.output_dir) / self.run_name

    # ── YAML I/O ────────────────────────────────────────────────────────

    @classmethod
    def from_yaml(cls, path: str | Path) -> "SDGConfig":
        with open(path) as f:
            data = yaml.safe_load(f) or {}
        valid_keys = {f.name for f in fields(cls)}
        filtered = {k: v for k, v in data.items() if k in valid_keys}
        config = cls(**filtered)
        if not config.mongo_uri:
            raise ValueError(
                f"mongo_uri is not set in {path}. "
                f"Add a line like `mongo_uri: \"mongodb://HOST:PORT\"` to the config. "
                f"(Common mistake: the key must be `mongo_uri` with an underscore, not `mongo-uri`.)"
            )
        return config

    def to_yaml(self, path: str | Path) -> None:
        Path(path).parent.mkdir(parents=True, exist_ok=True)
        with open(path, "w") as f:
            yaml.dump(asdict(self), f, default_flow_style=False, sort_keys=False)