Download sdg/config.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 3.17 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/config.py
- Command line
-
hf download hf://fzzhang/svd-code/sdg/config.py
-
curl -L -o config.py https://huggingface.co/fzzhang/svd-code/resolve/main/sdg/config.py
3.17 kB
| """ | |
| SDGConfig dataclass and YAML serialization. | |
| """ | |
| from dataclasses import dataclass, field, fields, asdict | |
| from pathlib import Path | |
| import yaml | |
| 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 ββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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 | |
| def mongo_db_name(self) -> str: | |
| return "sdg_cache" | |
| def output_path(self) -> Path: | |
| return Path(self.output_dir) / self.run_name | |
| # ββ YAML I/O ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| 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) | |