mirror of
https://github.com/wassname/AntiPaSTO.git
synced 2026-09-11 11:52:42 +08:00
Merge branch 'main' of https://github.com/wassname/AntiPaSTO
This commit is contained in:
@@ -142,3 +142,8 @@ Built on the shoulders of other chefs:
|
||||
url = {https://arxiv.org/abs/2601.07473}
|
||||
}
|
||||
```
|
||||
|
||||
|
||||
|
||||
<img width="832" height="554" alt="Nano banana's attempt to draw the loss landscape, I'm not sure if it helps understand the loss, but I like it" src="https://github.com/user-attachments/assets/1363cafe-24a9-458a-8f05-66b27cf5cb12" />
|
||||
|
||||
|
||||
+109
-32
@@ -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,
|
||||
),
|
||||
),
|
||||
|
||||
|
||||
@@ -42,6 +42,7 @@ from einops import einsum
|
||||
from jaxtyping import Float
|
||||
from loguru import logger
|
||||
from tqdm import tqdm
|
||||
import numpy as np
|
||||
|
||||
|
||||
def get_hidden_size(model: nn.Module) -> int:
|
||||
@@ -213,7 +214,7 @@ def approx_intersection_bases(
|
||||
V_a: Float[Tensor, "d r_a"],
|
||||
V_b: Float[Tensor, "d r_b"],
|
||||
top_k: int = 256,
|
||||
min_overlap: float = 0.1,
|
||||
min_overlap: float = 0.5,
|
||||
) -> tuple[Float[Tensor, "d k"], Float[Tensor, "k"]]:
|
||||
"""Intersection of two subspaces via principal angles.
|
||||
|
||||
@@ -229,9 +230,9 @@ def approx_intersection_bases(
|
||||
Args:
|
||||
V_a, V_b: Orthonormal bases [d_model, rank]
|
||||
top_k: Maximum number of intersection directions to return
|
||||
min_overlap: Minimum cos(principal_angle) to include (default 0.1).
|
||||
S=1 means perfect overlap, S=0 means orthogonal.
|
||||
Directions with S < min_overlap are excluded as "not truly shared".
|
||||
min_overlap: Minimum cos(principal_angle) to include (default 0.5 = 60°).
|
||||
S=1 means perfect overlap (0°), S=0 means orthogonal (90°).
|
||||
0.1 (cos 84°) was too permissive; 0.5 (cos 60°) ensures actual alignment.
|
||||
|
||||
Returns:
|
||||
V_shared: [d_model, k] orthonormal basis of shared directions
|
||||
@@ -253,8 +254,15 @@ def approx_intersection_bases(
|
||||
n_high_overlap = high_overlap_mask.sum().item()
|
||||
k = min(top_k, n_high_overlap, S.shape[0])
|
||||
|
||||
# Log overlap quality for diagnostics
|
||||
mean_overlap = S[:min(10, len(S))].mean().item()
|
||||
logger.debug(f"Intersection overlap: max={S[0]:.3f}, top10_mean={mean_overlap:.3f}, n>{min_overlap}={n_high_overlap}/{len(S)}")
|
||||
|
||||
if k == 0:
|
||||
logger.warning(f"intersect_bases: no directions with overlap > {min_overlap} (max S={S[0]:.3f}). Returning top-1 anyway.")
|
||||
logger.warning(
|
||||
f"Intersection: no directions with overlap > {min_overlap:.2f} (cos {min_overlap:.2f} = {np.arccos(min_overlap)*180/np.pi:.0f}°). "
|
||||
f"Max overlap={S[0]:.3f} ({np.arccos(S[0].item())*180/np.pi:.0f}°). Returning top-1 fallback."
|
||||
)
|
||||
k = 1
|
||||
|
||||
# Directions in original space:
|
||||
|
||||
@@ -1187,7 +1187,8 @@ def _compute_steering_f1_for_method(
|
||||
common_idx = df_neg.index.intersection(df_0.index).intersection(df_pos.index)
|
||||
if len(common_idx) == 0:
|
||||
return {"steering_f1": np.nan, "net_correct": np.nan, "correct_w": np.nan,
|
||||
"wrong_w": np.nan, "arb_w": np.nan, "precision": np.nan,
|
||||
"wrong_w": np.nan, "arb_w": np.nan, "correct_rate": np.nan,
|
||||
"wrong_rate": np.nan, "arb_rate": np.nan, "precision": np.nan,
|
||||
"recall": np.nan, "pmass_ratio": np.nan}
|
||||
|
||||
df_neg = df_neg.loc[common_idx]
|
||||
|
||||
@@ -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,
|
||||
@@ -1030,23 +1039,47 @@ def train_epoch(
|
||||
)
|
||||
wandb_run.log(val_metrics, step=step)
|
||||
|
||||
# Early stopping with min_delta (relative improvement threshold)
|
||||
# Early stopping (disabled when patience=0, e.g., with one-cycle scheduler)
|
||||
# Skip early stopping during warmup - LR is still ramping up
|
||||
# Early stopping with min_delta (relative improvement threshold).
|
||||
# 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
|
||||
in_warmup = opt_step < warmup_steps
|
||||
# Detect first validation AFTER warmup (best_val_loss still at inf means we haven't started tracking)
|
||||
first_post_warmup = (not in_warmup) and (best_val_loss[0] == float("inf"))
|
||||
|
||||
if in_warmup:
|
||||
logger.debug(f"Warmup: opt_step {opt_step}/{warmup_steps}, skipping early stopping check")
|
||||
elif first_post_warmup:
|
||||
# First validation after warmup - reset best_val_loss to current
|
||||
# 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
|
||||
|
||||
# "focus" is considered enabled unless we explicitly ignore it.
|
||||
focus_enabled = config.focus_softness < 1.0
|
||||
early_stop_enabled = (
|
||||
config.early_stop_patience > 0
|
||||
and config.coh
|
||||
and config.mono
|
||||
and focus_enabled
|
||||
and best_val_loss is not None
|
||||
and patience_counter is not None
|
||||
)
|
||||
early_stop_ready_step = max(warmup_steps, mono_warmup_steps, coh_warmup_steps)
|
||||
early_stop_ready = opt_step >= early_stop_ready_step
|
||||
first_post_ready = early_stop_enabled and early_stop_ready and (best_val_loss[0] == float("inf"))
|
||||
|
||||
if early_stop_enabled and not early_stop_ready:
|
||||
logger.debug(
|
||||
f"Early stop gated: opt_step {opt_step}/{early_stop_ready_step} "
|
||||
f"(warmup={warmup_steps}, mono_warmup={mono_warmup_steps}, coh_warmup={coh_warmup_steps})"
|
||||
)
|
||||
elif first_post_ready:
|
||||
best_val_loss[0] = val_loss
|
||||
patience_counter[0] = 0
|
||||
logger.info(f"Warmup complete at opt_step {opt_step}/{warmup_steps}. Starting early stopping with val_loss={val_loss:.4f}")
|
||||
|
||||
if config.early_stop_patience > 0 and best_val_loss is not None and patience_counter is not None and not in_warmup and not first_post_warmup:
|
||||
logger.info(
|
||||
f"Early stopping enabled at opt_step {opt_step}/{early_stop_ready_step}. "
|
||||
f"Starting tracking with val_loss={val_loss:.4f}"
|
||||
)
|
||||
|
||||
if early_stop_enabled and early_stop_ready and not first_post_ready:
|
||||
# Require relative improvement > min_delta to count as "better"
|
||||
improved = val_loss < best_val_loss[0] * (1 - config.early_stop_min_delta)
|
||||
|
||||
@@ -1826,6 +1859,7 @@ def train_model(config: TrainingConfig):
|
||||
num_workers=0 if config.quick else 8,
|
||||
pin_memory=True,
|
||||
persistent_workers=False if config.quick else True,
|
||||
drop_last=True, # need full batch for fisher
|
||||
)
|
||||
val_dataloader = DataLoader(
|
||||
val_dataset_pt,
|
||||
@@ -1835,17 +1869,25 @@ def train_model(config: TrainingConfig):
|
||||
num_workers=0 if config.quick else 8,
|
||||
pin_memory=True,
|
||||
persistent_workers=False if config.quick else True,
|
||||
drop_last=True, # need full batch for fisher
|
||||
)
|
||||
|
||||
total_steps = config.n_epochs * len(train_dataloader) // config.grad_accum_steps
|
||||
opt = torch.optim.AdamW(
|
||||
model.parameters(), lr=config.lr, weight_decay=config.wd
|
||||
)
|
||||
focus_enabled = config.focus_softness < 1.0
|
||||
early_stop_enabled = (
|
||||
config.early_stop_patience > 0 and config.coh and config.mono and focus_enabled
|
||||
)
|
||||
scheduler = torch.optim.lr_scheduler.OneCycleLR(
|
||||
opt, max_lr=config.lr, total_steps=total_steps, pct_start=config.warmup_pct,
|
||||
|
||||
# Early stopping and one cycle are not usually combined, this setting effectively turns it into constant LR with warmup
|
||||
final_div_factor=1.0 if (config.early_stop_patience > 0) else 1e5
|
||||
opt,
|
||||
max_lr=config.lr,
|
||||
total_steps=total_steps,
|
||||
pct_start=config.warmup_pct,
|
||||
# Early stopping and one-cycle are not usually combined; when early stopping
|
||||
# is enabled, use effectively-constant LR with warmup.
|
||||
final_div_factor=1.0 if early_stop_enabled else 1e5,
|
||||
)
|
||||
|
||||
logger.info(f"Training: {config.n_epochs} epochs, {total_steps} steps")
|
||||
|
||||
Reference in New Issue
Block a user