fix: on-policy data paths, 4-bit inference, revert adapter defaults

- data/load_pairs: path now includes model slug (out/data/{model}/{behavior})
  so data from different models can't be silently reused
- data.py, kl_calibrate.py, tinymfv_airisk.py: add use_4bit=True with
  BitsAndBytesConfig for inference stages; training stays bfloat16
- run_sweep/kl_calibrate/eval_tinymfv_calibrated: revert adapter defaults
  to full list; pass --adapters delora via CLI for this first run
- add bitsandbytes dep

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
wassname
2026-05-03 17:31:09 +08:00
co-authored by Claude Sonnet 4.6
parent 7396bc1544
commit 57a08750b8
8 changed files with 46 additions and 16 deletions
+1
View File
@@ -20,6 +20,7 @@ dependencies = [
"baukit @ git+https://github.com/davidbau/baukit.git",
"tiny-mfv @ git+https://github.com/wassname/tinymfv",
"flash-attn @ https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.6.3%2Bcu130torch2.11-cp311-cp311-linux_x86_64.whl ; python_version == '3.11' and sys_platform == 'linux'",
"bitsandbytes>=0.49.2",
]
[project.optional-dependencies]
+10 -5
View File
@@ -27,7 +27,7 @@ import tyro
from datasets import Dataset
from loguru import logger
from tqdm.auto import tqdm
from transformers import AutoModelForCausalLM, AutoTokenizer, StaticCache
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, StaticCache
from ws._log import get_argv, setup_logging
from ws._tok_extras import chat_template_extras, has_thinking_mode
@@ -258,6 +258,7 @@ class DataCfg:
n_personas: int = 5
n_samples: int = 10
out: Path = Path("out/data")
use_4bit: bool = True
batch_size: int = 8
min_new_tokens: int = 1024
max_new_tokens: int = 1280
@@ -543,8 +544,10 @@ def generate_pairs(cfg: DataCfg) -> Path:
tok = AutoTokenizer.from_pretrained(cfg.model_id)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
bnb_cfg = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16) if cfg.use_4bit else None
model = AutoModelForCausalLM.from_pretrained(
cfg.model_id, torch_dtype=torch.bfloat16, device_map="cuda"
cfg.model_id, torch_dtype=torch.bfloat16, device_map="cuda",
quantization_config=bnb_cfg,
)
model.eval()
@@ -613,15 +616,17 @@ def generate_pairs(cfg: DataCfg) -> Path:
ds = Dataset.from_list(rows)
assert_generated_pairs_diverged(ds)
out_dir = cfg.out / cfg.behavior
model_slug = cfg.model_id.replace("/", "_")
out_dir = cfg.out / model_slug / cfg.behavior
out_dir.mkdir(parents=True, exist_ok=True)
ds.save_to_disk(str(out_dir))
logger.info(f"saved {len(ds)} pairs to {out_dir}")
return out_dir
def load_pairs(behavior: str, root: Path = Path("out/data")) -> Dataset:
ds = Dataset.load_from_disk(str(root / behavior))
def load_pairs(behavior: str, model_id: str, root: Path = Path("out/data")) -> Dataset:
model_slug = model_id.replace("/", "_")
ds = Dataset.load_from_disk(str(root / model_slug / behavior))
assert_generated_pairs_diverged(ds)
return ds
+4 -2
View File
@@ -20,7 +20,7 @@ import tyro
from datasets import load_dataset
from loguru import logger
from tabulate import tabulate
from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from ws._artifacts import model_slug, timestamp_prefix
from ws._log import final_summary, get_argv, setup_logging
@@ -70,6 +70,7 @@ class TinyMFVAiriskCfg:
batch_size: int = 16
max_length: int = 256
limit: int = 0
use_4bit: bool = True
bootstrap_samples: int = 1000
bootstrap_seed: int = 0
@@ -520,7 +521,8 @@ def run_eval(cfg: TinyMFVAiriskCfg) -> tuple[pl.DataFrame, pl.DataFrame, pl.Data
if tok.pad_token is None:
tok.pad_token = tok.eos_token
tok.padding_side = "left"
model = AutoModelForCausalLM.from_pretrained(cfg.model, torch_dtype=torch.bfloat16, device_map="cuda")
bnb_cfg = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16) if cfg.use_4bit else None
model = AutoModelForCausalLM.from_pretrained(cfg.model, torch_dtype=torch.bfloat16, device_map="cuda", quantization_config=bnb_cfg)
model.eval()
vignettes = _load_vignettes(cfg.limit)
+6 -3
View File
@@ -42,7 +42,7 @@ import tyro
from loguru import logger
from tabulate import tabulate
from torch import Tensor
from transformers import AutoModelForCausalLM, AutoTokenizer
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from ws._log import final_summary, get_argv, setup_logging
from ws.data import _load_suffixes
@@ -66,7 +66,7 @@ class KLCalibrateCfg:
model: str = "Qwen/Qwen3-0.6B"
behavior: str = "honesty"
out: Path = Path("out")
adapters: tuple[str, ...] = ("delora",)
adapters: tuple[str, ...] = ("lora", "pissa", "dora", "delora", "oft", "ia3")
n_calib_prompts: int = 50
n_audit_prompts: int = 100
n_tokens: int = 50
@@ -80,6 +80,7 @@ class KLCalibrateCfg:
bracket_hi: float = 16.0
n_root_iters: int = 12 # Illinois inner loop; usually converges in 3-5
convergence_tol: float = 0.05 # |p95 - target| < tol (absolute, in nats)
use_4bit: bool = True
seed: int = 0
@@ -322,8 +323,10 @@ def main(cfg: KLCalibrateCfg) -> None:
if tok.pad_token is None:
tok.pad_token = tok.eos_token
tok.padding_side = "left"
bnb_cfg = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_compute_dtype=torch.bfloat16) if cfg.use_4bit else None
model = AutoModelForCausalLM.from_pretrained(
cfg.model, torch_dtype=torch.bfloat16, device_map="cuda"
cfg.model, torch_dtype=torch.bfloat16, device_map="cuda",
quantization_config=bnb_cfg,
)
model.eval()
+4 -3
View File
@@ -50,12 +50,13 @@ class Cfg:
def _maybe_data(cfg: Cfg) -> Dataset:
from ws.data import _personas
data_root = cfg.data_root
behavior_dir = data_root / cfg.behavior
model_slug = cfg.model.replace("/", "_")
behavior_dir = data_root / model_slug / cfg.behavior
sys_pos_all, _ = _personas(cfg.behavior)
n_personas = min(cfg.n_personas, len(sys_pos_all))
expected = cfg.n_topics * n_personas * cfg.n_samples
if behavior_dir.exists():
ds = load_pairs(cfg.behavior, root=data_root)
ds = load_pairs(cfg.behavior, cfg.model, root=data_root)
if len(ds) != expected:
raise ValueError(
f"on-disk data at {behavior_dir} has {len(ds)} pairs but "
@@ -77,7 +78,7 @@ def _maybe_data(cfg: Cfg) -> Dataset:
presence_penalty=cfg.data_presence_penalty,
)
generate_pairs(dcfg)
return load_pairs(cfg.behavior, root=data_root)
return load_pairs(cfg.behavior, cfg.model, root=data_root)
def main(cfg: Cfg) -> None:
+1 -1
View File
@@ -24,7 +24,7 @@ from ws.replicate import main as replicate_main
class SweepCfg:
model: str = "Qwen/Qwen3-0.6B"
behavior: str = "authority"
adapters: tuple[str, ...] = ("delora",)
adapters: tuple[str, ...] = ("lora", "dora", "pissa", "delora", "oft", "boft", "ia3")
rank: int = 32
lr: float = 2e-4
epochs: float = 1.0
+1 -1
View File
@@ -28,7 +28,7 @@ from loguru import logger
class EvalTinymfvCalibratedCfg:
behavior: str = "authority"
out: Path = Path("out")
adapters: tuple[str, ...] = ("delora",)
adapters: tuple[str, ...] = ("lora", "dora", "pissa", "delora", "oft", "ia3")
model: str = "Qwen/Qwen3.5-4B"
bootstrap_samples: int = 256
limit: int = 0
Generated
+19 -1
View File
@@ -14,7 +14,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-04-28T08:51:29.029061397Z"
exclude-newer = "2026-04-28T09:29:10.855326898Z"
exclude-newer-span = "P5D"
[[package]]
@@ -236,6 +236,22 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/71/cc/18245721fa7747065ab478316c7fea7c74777d07f37ae60db2e84f8172e8/beartype-0.22.9-py3-none-any.whl", hash = "sha256:d16c9bbc61ea14637596c5f6fbff2ee99cbe3573e46a716401734ef50c3060c2", size = 1333658, upload-time = "2025-12-13T06:50:28.266Z" },
]
[[package]]
name = "bitsandbytes"
version = "0.49.2"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "numpy" },
{ name = "packaging" },
{ name = "torch" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/d8/7d/f1fe0992334b18cd8494f89aeec1dcc674635584fcd9f115784fea3a1d05/bitsandbytes-0.49.2-py3-none-macosx_14_0_arm64.whl", hash = "sha256:87be5975edeac5396d699ecbc39dfc47cf2c026daaf2d5852a94368611a6823f", size = 131940, upload-time = "2026-02-16T21:26:04.572Z" },
{ url = "https://files.pythonhosted.org/packages/29/71/acff7af06c818664aa87ff73e17a52c7788ad746b72aea09d3cb8e424348/bitsandbytes-0.49.2-py3-none-manylinux_2_24_aarch64.whl", hash = "sha256:2fc0830c5f7169be36e60e11f2be067c8f812dfcb829801a8703735842450750", size = 31442815, upload-time = "2026-02-16T21:26:06.783Z" },
{ url = "https://files.pythonhosted.org/packages/19/57/3443d6f183436fbdaf5000aac332c4d5ddb056665d459244a5608e98ae92/bitsandbytes-0.49.2-py3-none-manylinux_2_24_x86_64.whl", hash = "sha256:54b771f06e1a3c73af5c7f16ccf0fc23a846052813d4b008d10cb6e017dd1c8c", size = 60651714, upload-time = "2026-02-16T21:26:11.579Z" },
{ url = "https://files.pythonhosted.org/packages/b6/d4/501655842ad6771fb077f576d78cbedb5445d15b1c3c91343ed58ca46f0e/bitsandbytes-0.49.2-py3-none-win_amd64.whl", hash = "sha256:2e0ddd09cd778155388023cbe81f00afbb7c000c214caef3ce83386e7144df7d", size = 55372289, upload-time = "2026-02-16T21:26:16.267Z" },
]
[[package]]
name = "certifi"
version = "2026.2.25"
@@ -3028,6 +3044,7 @@ dependencies = [
{ name = "accelerate" },
{ name = "baukit" },
{ name = "beartype" },
{ name = "bitsandbytes" },
{ name = "datasets" },
{ name = "einops" },
{ name = "flash-attn", marker = "python_full_version < '3.12' and sys_platform == 'linux'" },
@@ -3055,6 +3072,7 @@ requires-dist = [
{ name = "accelerate", specifier = ">=1.0" },
{ name = "baukit", git = "https://github.com/davidbau/baukit.git" },
{ name = "beartype", specifier = ">=0.19" },
{ name = "bitsandbytes", specifier = ">=0.49.2" },
{ name = "datasets", specifier = ">=3.0" },
{ name = "einops", specifier = ">=0.8" },
{ name = "flash-attn", marker = "python_full_version == '3.11.*' and sys_platform == 'linux'", url = "https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.9.4/flash_attn-2.6.3%2Bcu130torch2.11-cp311-cp311-linux_x86_64.whl" },