mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-18 00:47:23 +02:00
178 lines
5.8 KiB
Python
178 lines
5.8 KiB
Python
"""YAML-based configuration for ablation runs."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import yaml
|
|
|
|
|
|
@dataclass
|
|
class ModelConfig:
|
|
name: str
|
|
task: str = "causal_lm"
|
|
dtype: str = "float32"
|
|
device: str = "auto"
|
|
trust_remote_code: bool = False
|
|
num_labels: int = 2
|
|
quantization: str | None = None
|
|
|
|
|
|
@dataclass
|
|
class DatasetConfig:
|
|
name: str
|
|
split: str = "test"
|
|
subset: str | None = None
|
|
text_column: str = "text"
|
|
label_column: str = "label"
|
|
max_samples: int | None = None
|
|
|
|
|
|
@dataclass
|
|
class StrategyConfig:
|
|
name: str
|
|
params: dict[str, Any] = field(default_factory=dict)
|
|
|
|
|
|
@dataclass
|
|
class RemoteConfig:
|
|
"""Optional remote execution settings for running on a GPU node via SSH."""
|
|
|
|
host: str
|
|
user: str = "root"
|
|
port: int = 22
|
|
ssh_key: str | None = None
|
|
known_hosts_file: str | None = None
|
|
remote_dir: str = "/tmp/obliteratus_run"
|
|
install_timeout: int = 600
|
|
python: str = "python3"
|
|
sync_results: bool = True
|
|
gpus: str | None = None # comma-separated GPU IDs or "all"
|
|
install_source: str = "git+https://github.com/elder-plinius/OBLITERATUS.git"
|
|
|
|
def __post_init__(self) -> None:
|
|
from obliteratus.remote_contracts import validate_remote_settings
|
|
|
|
self.gpus = validate_remote_settings(
|
|
host=self.host,
|
|
user=self.user,
|
|
port=self.port,
|
|
remote_dir=self.remote_dir,
|
|
python=self.python,
|
|
gpus=self.gpus,
|
|
install_source=self.install_source,
|
|
)
|
|
if (
|
|
isinstance(self.install_timeout, bool)
|
|
or not isinstance(self.install_timeout, int)
|
|
or self.install_timeout <= 0
|
|
):
|
|
raise ValueError("remote install timeout must be a positive integer")
|
|
|
|
|
|
@dataclass
|
|
class StudyConfig:
|
|
"""Top-level configuration for an ablation run."""
|
|
|
|
model: ModelConfig
|
|
dataset: DatasetConfig
|
|
strategies: list[StrategyConfig]
|
|
metrics: list[str] = field(default_factory=lambda: ["perplexity"])
|
|
batch_size: int = 8
|
|
max_length: int = 512
|
|
output_dir: str = "results"
|
|
remote: RemoteConfig | None = None
|
|
|
|
@classmethod
|
|
def from_yaml(cls, path: str | Path) -> StudyConfig:
|
|
path = Path(path)
|
|
raw = yaml.safe_load(path.read_text())
|
|
return cls.from_dict(raw)
|
|
|
|
@classmethod
|
|
def from_dict(cls, d: dict) -> StudyConfig:
|
|
# Accept both "preset" and legacy "study_preset" keys.
|
|
if "preset" in d and "study_preset" not in d:
|
|
d["study_preset"] = d["preset"]
|
|
# If a study_preset key is provided, use it as the base and allow
|
|
# the rest of the config to override individual fields.
|
|
if "study_preset" in d:
|
|
from obliteratus.study_presets import get_study_preset
|
|
|
|
preset = get_study_preset(d["study_preset"])
|
|
# Preset provides defaults; explicit keys in the dict override.
|
|
if "strategies" not in d:
|
|
d["strategies"] = preset.strategies
|
|
if "metrics" not in d:
|
|
d["metrics"] = preset.metrics
|
|
if "batch_size" not in d:
|
|
d["batch_size"] = preset.batch_size
|
|
if "max_length" not in d:
|
|
d["max_length"] = preset.max_length
|
|
# Preset max_samples flows into dataset if not set
|
|
ds = d.get("dataset", {})
|
|
if "max_samples" not in ds and ds:
|
|
ds["max_samples"] = preset.max_samples
|
|
d["dataset"] = ds
|
|
|
|
model = ModelConfig(**d["model"])
|
|
dataset = DatasetConfig(**d["dataset"])
|
|
strategies = [StrategyConfig(**s) for s in d["strategies"]]
|
|
remote = None
|
|
if "remote" in d and d["remote"]:
|
|
remote = RemoteConfig(**d["remote"])
|
|
|
|
return cls(
|
|
model=model,
|
|
dataset=dataset,
|
|
strategies=strategies,
|
|
metrics=d.get("metrics", ["perplexity"]),
|
|
batch_size=d.get("batch_size", 8),
|
|
max_length=d.get("max_length", 512),
|
|
output_dir=d.get("output_dir", "results"),
|
|
remote=remote,
|
|
)
|
|
|
|
def to_dict(self) -> dict:
|
|
result = {
|
|
"model": {
|
|
"name": self.model.name,
|
|
"task": self.model.task,
|
|
"dtype": self.model.dtype,
|
|
"device": self.model.device,
|
|
"trust_remote_code": self.model.trust_remote_code,
|
|
"num_labels": self.model.num_labels,
|
|
"quantization": self.model.quantization,
|
|
},
|
|
"dataset": {
|
|
"name": self.dataset.name,
|
|
"split": self.dataset.split,
|
|
"subset": self.dataset.subset,
|
|
"text_column": self.dataset.text_column,
|
|
"label_column": self.dataset.label_column,
|
|
"max_samples": self.dataset.max_samples,
|
|
},
|
|
"strategies": [{"name": s.name, "params": s.params} for s in self.strategies],
|
|
"metrics": self.metrics,
|
|
"batch_size": self.batch_size,
|
|
"max_length": self.max_length,
|
|
"output_dir": self.output_dir,
|
|
}
|
|
if self.remote is not None:
|
|
result["remote"] = {
|
|
"host": self.remote.host,
|
|
"user": self.remote.user,
|
|
"port": self.remote.port,
|
|
"ssh_key": self.remote.ssh_key,
|
|
"known_hosts_file": self.remote.known_hosts_file,
|
|
"remote_dir": self.remote.remote_dir,
|
|
"install_timeout": self.remote.install_timeout,
|
|
"python": self.remote.python,
|
|
"sync_results": self.remote.sync_results,
|
|
"gpus": self.remote.gpus,
|
|
"install_source": self.remote.install_source,
|
|
}
|
|
return result
|