""" 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)