mirror of
https://github.com/wassname/weight-steering.git
synced 2026-09-12 13:00:32 +08:00
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:
co-authored by
Claude Sonnet 4.6
parent
7396bc1544
commit
57a08750b8
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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" },
|
||||
|
||||
Reference in New Issue
Block a user