svd-code / sdg /config.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
3.17 kB
"""
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)