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