danielfein's picture
Add training support package
a4019dd verified
Raw
History Blame Contribute Delete
4.88 kB
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
def model_slug(model_name: str) -> str:
return model_name.replace("/", "__")
@dataclass(slots=True)
class PangramBinaryConfig:
enabled: bool = True
dataset_name: str = "pangram/editlens_iclr"
dataset_split: str = "train"
local_dataset_path: Path | None = Path("/home/ubuntu/data/pangram_editlens_iclr")
ai_text_types: tuple[str, ...] = ("ai_generated",)
human_text_types: tuple[str, ...] = ("human_written",)
train_pairs: int = 5_000
@dataclass(slots=True)
class RaidBinaryConfig:
enabled: bool = True
dataset_name: str = "liamdugan/raid"
dataset_split: str = "train"
human_model_name: str = "human"
require_attack_none: bool = True
train_pairs: int = 5_000
eval_holdout_pairs: int = 1_000
@dataclass(slots=True)
class DataConfig:
task_name: str = "binary_human_vs_ai"
training_holdout_pairs: int = 1_000
min_text_chars: int = 200
pangram: PangramBinaryConfig = field(default_factory=PangramBinaryConfig)
raid: RaidBinaryConfig = field(default_factory=RaidBinaryConfig)
@dataclass(slots=True)
class ModelConfig:
model_name: str = "meta-llama/Llama-3.1-8B-Instruct"
ai_token: str = "<ai>"
human_token: str = "<human>"
max_length: int = 384
prompt_template: str = "Write {token} text."
@dataclass(slots=True)
class TrainingConfig:
ai_learning_rate: float = 1.0e-4
human_learning_rate: float = 1.0e-4
batch_size: int = 16
train_steps: int | None = 625
beta: float = 0.1
apo_alpha: float = 1.0
warmup_steps: int = 20
min_learning_rate: float = 1.0e-6
ref_cache_batch_size: int = 16
eval_subset_size: int = 128
eval_every_steps: int = 50
neutral_word: str | None = None # if set, warm-start new token rows from this vocab token
source_balance_by_source: bool = False
@dataclass(slots=True)
class CheckpointInitConfig:
ai_token_path: Path | None = None
human_token_path: Path | None = None
@dataclass(slots=True)
class EvaluationConfig:
mode: str = "holdout"
saved_pairs_filename: str = "holdout_pairs.json"
binary_split: str = "test"
positive_text_types: tuple[str, ...] = ("ai_generated",)
negative_text_types: tuple[str, ...] = ("human_written",)
binary_output_path: Path | None = None
@dataclass(slots=True)
class ScoringConfig:
text: str = "Example text to score."
score_mode: str = "avg_margin"
token_sigmoid_tau: float = 1.0
@dataclass(slots=True)
class VerbalizationTokenSetConfig:
name: str
token_dir: Path
def default_verbalization_token_sets() -> tuple[VerbalizationTokenSetConfig, ...]:
root = Path("/lambda/nfs/daniel-dc/detectiontokens_outputs")
return (
VerbalizationTokenSetConfig(
name="raid_only",
token_dir=root / "raid_only_train_raid_holdout_eval" / "tokens",
),
VerbalizationTokenSetConfig(
name="pangram_only",
token_dir=root / "pangram_1k_train_pangram_test_eval" / "tokens",
),
)
@dataclass(slots=True)
class VerbalizationConfig:
output_dir: Path = Path("/lambda/nfs/daniel-dc/detectiontokens_outputs/verbalizations")
token_sets: tuple[VerbalizationTokenSetConfig, ...] = field(default_factory=default_verbalization_token_sets)
n_samples: int = 5
max_new_tokens: int = 128
@dataclass(slots=True)
class OutputConfig:
output_root: Path = Path("/lambda/nfs/daniel-dc/detectiontokens_outputs/pangram_raid_binary_llama31_8b")
repo_root: Path = Path("/lambda/nfs/daniel-dc/DetectionTokens")
@property
def splits_dir(self) -> Path:
return self.output_root / "splits"
@property
def tokens_dir(self) -> Path:
return self.output_root / "tokens"
@property
def model_tokens_root(self) -> Path:
return self.repo_root / "tokens"
def model_tokens_dir(self, model_name: str) -> Path:
return self.model_tokens_root / model_slug(model_name)
@property
def cache_dir(self) -> Path:
return self.output_root / "ref_cache"
@property
def evaluation_dir(self) -> Path:
return self.output_root / "evaluation"
@dataclass(slots=True)
class PipelineConfig:
seed: int = 42
data: DataConfig = field(default_factory=DataConfig)
model: ModelConfig = field(default_factory=ModelConfig)
training: TrainingConfig = field(default_factory=TrainingConfig)
init_checkpoints: CheckpointInitConfig = field(default_factory=CheckpointInitConfig)
evaluation: EvaluationConfig = field(default_factory=EvaluationConfig)
scoring: ScoringConfig = field(default_factory=ScoringConfig)
verbalization: VerbalizationConfig = field(default_factory=VerbalizationConfig)
output: OutputConfig = field(default_factory=OutputConfig)