diff --git a/antipasto/config.py b/antipasto/config.py index 2db9036..b2e2ceb 100644 --- a/antipasto/config.py +++ b/antipasto/config.py @@ -40,7 +40,7 @@ class TrainingConfig: seed: int = 42 """Random seed for reproducibility (layer selection, dim selection, training dynamics).""" - init_n_samples: int = 1000 + init_n_samples: int = 2000 """Number of samples for WANDA-style dimension selection and subspace computation. Higher = more stable activation statistics, but slower init. @@ -52,10 +52,10 @@ class TrainingConfig: data_seed: int = 42 """Fixed seed for data selection (which suffixes are used).""" - model_name: str = "Qwen/Qwen3-4B-Instruct-2507" + model_name: str = "google/gemma-3-12b-it" quantization_type: Literal["4bit", "8bit", "none"] = "none" - n_modules: int = 512 + n_modules: int = 256 """Total number of layer×module combinations to select (by gradient importance). Examples with n_modules=5: @@ -64,7 +64,7 @@ class TrainingConfig: The selection is sparse: each layer×module is ranked by gradient, top-k are selected. Not Cartesian product - can have multiple modules at one layer, none at another. - Default 42 (≈14 layers × 3 modules for typical 36-layer model). + Default 256. """ target_modules: List[str] = ["residual-writers"] @@ -80,23 +80,25 @@ class TrainingConfig: Explicit list: ["down_proj", "o_proj"] - only these module suffixes are candidates. """ - bs: int = 8 + bs: int = 14 """Batch size""" - n_epochs: int = 20 + n_epochs: int = 30 - lr: float = 5e-4 - """Learning rate. Sweep findings: 1e-4 too low (-18 F1), 3e-5 too low (-17 F1). - For Cayley: 4e-4 to 6e-4. For LoRA/DoRA: ~10x lower (3e-5 to 6e-5). + lr: float = 0.002 + """Learning rate. + + Empirically, Cayley rotations tend to tolerate higher LR than LoRA/DoRA. + This repo's default matches a strong run on Qwen3-14B. """ - wd: float = 1e-8 + wd: float = 1e-4 """Weight decay""" n_logs: int = 10 """Log this many times per training""" - val_every_n_samples: int = 512 + val_every_n_samples: int = 1024 """Validate every N training samples (independent of logging).""" effective_bs: int = 32 @@ -108,7 +110,7 @@ class TrainingConfig: val_split: float = 0.15 """Fraction of data for validation""" - early_stop_patience: int = 11 + early_stop_patience: int = 16 """Stop if val loss doesn't improve for N validation checks. 0 = disabled (recommended with one-cycle scheduler).""" early_stop_min_delta: float = 0.00001 @@ -171,7 +173,7 @@ class TrainingConfig: - taskdiff_x_write_x_notlogits: Task ∩ write ∩ (lm_head^⊥) """ - loss_subspace_rank: Optional[int] = 8 + loss_subspace_rank: Optional[int] = 4 """Rank (top-k) for loss subspace. If None (default), select rank automatically via `loss_subspace_energy_frac` @@ -187,12 +189,12 @@ class TrainingConfig: k such that cumulative energy >= this fraction. 60% was used in MSRS paper """ - loss_layer_frac: float = 0.9 + loss_layer_frac: float = 0.5 """Depth fraction (0-1) at which to apply representation loss. The loss is computed at a single layer: int(loss_layer_frac * num_hidden_layers). - Default 0.8 (80% depth) is in the "planning zone" where Fisher ratio and + Default 0.5 (50% depth) is a simple mid-depth choice; prior sweeps often found a "planning zone" where Fisher ratio and cross-sample consistency peak across tested architectures (Qwen, Gemma). Also suppurted by supported by e.g 2024-Gurnee-Universal-Neurons-in-GPT2-Language-Models.md Rationale: gradient-based layer selection was circular (gradients flow FROM @@ -220,7 +222,7 @@ class TrainingConfig: dataset_name: str = "honest" - max_samples: Optional[int] = 800 + max_samples: Optional[int] = 3000 """Max training samples (None = all)""" n_last_tokens: int = 3 @@ -274,15 +276,15 @@ class TrainingConfig: Projection loss naturally creates ordering; mono is a safety rail, not driver. """ - mono_margin: float = 0.4 + mono_margin: float = 0.5 """Monotonic threshold_frac: fraction of √H_ref for minimum separation. Threshold = threshold_frac × √H_ref + threshold_floor. - With H_ref=4 nats (typical), threshold_frac=0.4, floor=0.04: threshold ≈ 0.84 nats. + With H_ref=4 nats (typical), threshold_frac=0.5, floor=0.04: threshold ≈ 1.04 nats. Sweep findings (2026-01-07, gemma1b): | margin | F1 | - | 0.4 | 23.6 | ← current default + | 0.5 | (default) | 0.2 | 17.2 | | 0.25 | 0.0 | (collapsed) @@ -300,7 +302,7 @@ class TrainingConfig: Prevents division issues and provides small stable deadzone. """ - mono_weight: float = 20.0 + mono_weight: float = 30.0 """Monotonicity loss scaling. WARNING: Values ≥100 trap adapters in bad init - can't learn "no change" at c=0. @@ -310,35 +312,36 @@ class TrainingConfig: """ mono_warmup_frac: float = -2 - """Constraint warmup using -N syntax (binary: off during warmup, on after). + """Constraint warmup using -N syntax (gradual linear ramp from 0 to full weight). - Negative (-N): N × warmup_pct (e.g., -2 = 2× LR warmup = 20% at default) - - Zero: No warmup, active from start + - Zero: No warmup, full weight from start - Positive: Explicit fraction (e.g., 0.3 = 30% of training) - Default -2: constraints activate at 2× LR warmup. Lets projection establish - direction before constraints kick in. + Default -2: mono ramps up over 2× LR warmup. Lets projection establish + direction before constraints reach full strength. """ coh_warmup_frac: float = -2 - """Constraint warmup using -N syntax (binary: off during warmup, on after). + """Constraint warmup using -N syntax (gradual linear ramp from 0 to full weight). - Negative (-N): N × warmup_pct (e.g., -2 = 2× LR warmup = 20% at default) - - Zero: No warmup, active from start + - Zero: No warmup, full weight from start - Positive: Explicit fraction (e.g., 0.3 = 30% of training) - Default -2: synchronized with mono_warmup_frac. + Default -2: coh ramps up over 2× LR warmup. Prevents coh from fighting + projection loss early when deltas are large. """ conc_warmup_frac: float = -2 - """Concentration weighting warmup using -N syntax (binary: off during warmup, on after). + """Concentration (focus) weighting warmup using -N syntax (gradual linear ramp). - Negative (-N): N × warmup_pct (e.g., -2 = 2× LR warmup = 20% at default) - - Zero: No warmup, active from start + - Zero: No warmup, full weight from start - Positive: Explicit fraction (e.g., 0.3 = 30% of training) - During warmup, delta_*_norm_full=None disables subspace focus weighting. - Default -2: synchronized with other constraints. + During ramp-up, focus weighting interpolates from 1.0 (no penalty for out-of-subspace) + to the configured focus_softness. Default -2: synchronized with other constraints. """ orth_weight: float = 0 @@ -647,12 +650,86 @@ default_configs = { bs=64, ), ), + # Qwen/Qwen3-32B + "q32b-80gb": ( + "Qwen 32B on 80GB GPU (maximum size)", + TrainingConfig( + model_name="Qwen/Qwen3-32B", + bs=4, + ), + ), + "q14b-80gb": ( + "Qwen 14B on 80GB GPU (production quality)", + TrainingConfig( + model_name="Qwen/Qwen3-14B", + bs=12, + ), + ), + + "q14b-goodrun": ( + "Qwen 14B best run (fisher, coh+mono, r64)", + TrainingConfig( + model_name="Qwen/Qwen3-14B", + quantization_type="none", + n_modules=256, + target_modules=["residual-writers"], + bs=12, + n_epochs=20, + lr=0.002, + wd=1e-4, + n_logs=10, + val_every_n_samples=512, + effective_bs=32, + quick=False, + val_split=0.15, + early_stop_patience=14, + early_stop_min_delta=1e-5, + warmup_pct=0.1, + r=64, + rot_u=False, + rot_v=True, + dim_select_method="wanda_svd_l1_trip", + max_rotation_angle=pi / 4, + loss_subspace="taskdiff_x_suppressed_x_write", + loss_subspace_rank=4, + loss_subspace_energy_frac=0.6, + loss_layer_frac=0.5, + min_adapter_layer_frac=0.1, + dataset_name="honest", + max_samples=3000, + n_last_tokens=3, + coh=True, + coh_weight=10.0, + coh_thresh=0.9, + coh_barrier_mode="log1p_squared", + coh_lse_temperature=3.0, + mono=True, + mono_margin=0.5, + mono_threshold_floor=0.04, + mono_weight=30.0, + mono_warmup_frac=-2, + coh_warmup_frac=-2, + orth_weight=0, + antisym_margin=0.0, + fisher_var_floor_frac=0.1, + fisher_abs_std_floor=0.05, + fisher_detach_std=True, + eval_max_dilemmas=None, + eval_max_tokens=288, + use_wandb=True, + wandb_project="AntiPaSTO", + wandb_tags=None, + verbose=1, + PROMPT=PROMPT, + PERSONAS=PERSONAS, + ), + ), # add gemma4b "gemma12b-80gb": ( "Gemma 3 12B on 80GB GPU", TrainingConfig( model_name="google/gemma-3-12b-it", - bs=4, + bs=14, ), ), diff --git a/antipasto/train/train_adapter.py b/antipasto/train/train_adapter.py index 0121f93..5a6dda3 100644 --- a/antipasto/train/train_adapter.py +++ b/antipasto/train/train_adapter.py @@ -126,10 +126,19 @@ def compute_batch_loss( coh_warmup_steps = resolve_warmup(config.coh_warmup_frac) conc_warmup_steps = resolve_warmup(config.conc_warmup_frac) - # Binary switch: constraints off during warmup, on after - effective_mono_weight = config.mono_weight if step >= mono_warmup_steps else 0.0 - enable_coherence_effective = config.coh and (step >= coh_warmup_steps) - enable_concentration = step >= conc_warmup_steps + # Gradual linear ramp: 0 → full weight over warmup period + def ramp_weight(weight: float, warmup_steps: int) -> float: + if warmup_steps <= 0: + return weight + progress = min(1.0, step / warmup_steps) + return weight * progress + + effective_mono_weight = ramp_weight(config.mono_weight, mono_warmup_steps) + effective_coh_weight = ramp_weight(config.coh_weight, coh_warmup_steps) if config.coh else 0.0 + enable_coherence_effective = config.coh # Always enabled if config.coh=True, weight controls strength + # Focus warmup: interpolate focus_softness from 1.0 (disabled) to configured value + focus_warmup_progress = min(1.0, step / conc_warmup_steps) if conc_warmup_steps > 0 else 1.0 + effective_focus_softness = 1.0 + (config.focus_softness - 1.0) * focus_warmup_progress attention_mask = batch["attention_mask"] mask_cho = attention_mask[::2] @@ -258,7 +267,7 @@ def compute_batch_loss( delta_neg_norm_full = delta_neg_agg.norm(dim=-1) # [b] # Antisymmetric loss (Fisher + align + delta_full) - # Disable concentration during warmup (delta_norm_full=None) + # Focus now uses gradual ramp via effective_focus_softness (1.0 = disabled, <1 = enabled) loss_dict = contrastive_steering_loss_with_ref( s_ref_cho=s_ref_cho, s_ref_rej=s_ref_rej, @@ -270,9 +279,9 @@ def compute_batch_loss( last_n_tokens=config.n_last_tokens, orth_weight=config.orth_weight, antisym_margin=config.antisym_margin, - focus_softness=config.focus_softness, - delta_pos_norm_full=delta_pos_norm_full if enable_concentration else None, - delta_neg_norm_full=delta_neg_norm_full if enable_concentration else None, + focus_softness=effective_focus_softness, + delta_pos_norm_full=delta_pos_norm_full, + delta_neg_norm_full=delta_neg_norm_full, fisher_var_floor_frac=config.fisher_var_floor_frac, fisher_abs_std_floor=config.fisher_abs_std_floor, fisher_detach_std=config.fisher_detach_std, @@ -324,7 +333,7 @@ def compute_batch_loss( ref_label_logp=ref_coherence, pi_label_logp=pi_coherence, mask=mask_logp, - scale=config.coh_weight, + scale=effective_coh_weight, ref_logits=ref_logits, pi_logits=pi_logits, coh_thresh_frac=config.coh_thresh, @@ -1034,8 +1043,12 @@ def train_epoch( # Enable ONLY when coherence + monotonic + focus are ON, and only after # all warmups are finished (LR warmup + coh/mono warmups). warmup_steps = int(total_steps * config.warmup_pct) if total_steps else 0 - mono_warmup_frac = config.mono_warmup_frac if config.mono_warmup_frac >= 0 else config.warmup_pct - coh_warmup_frac = config.coh_warmup_frac if config.coh_warmup_frac >= 0 else config.warmup_pct + + # resolve_warmup: -N → N × warmup_pct, else explicit + def resolve_warmup_frac(frac: float) -> float: + return (-frac) * config.warmup_pct if frac < 0 else frac + mono_warmup_frac = resolve_warmup_frac(config.mono_warmup_frac) + coh_warmup_frac = resolve_warmup_frac(config.coh_warmup_frac) mono_warmup_steps = int(total_steps * mono_warmup_frac) if total_steps else 0 coh_warmup_steps = int(total_steps * coh_warmup_frac) if total_steps else 0