benjamin5607 commited on
Commit
2aa5051
·
verified ·
1 Parent(s): da2830e

Add safety_eval deps for Space

Browse files
Files changed (1) hide show
  1. safety_eval/platform/local_model.py +524 -0
safety_eval/platform/local_model.py ADDED
@@ -0,0 +1,524 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load and run the fine-tuned Jekyll & Hyde model (dual LoRA or merged weights)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import json
6
+ import re
7
+ from dataclasses import dataclass
8
+ from pathlib import Path
9
+ from typing import Any, Literal
10
+
11
+ import yaml
12
+
13
+ ROOT = Path(__file__).resolve().parent.parent.parent
14
+ MERGED_DIR = ROOT / "models" / "merged" / "jekyll-hyde"
15
+ MANIFEST_PATH = MERGED_DIR / "jekyll_hyde_manifest.json"
16
+ TRAIN_CONFIG = ROOT / "training" / "config.yaml"
17
+
18
+ AdapterName = Literal["jekyll", "hyde"]
19
+
20
+ _model = None
21
+ _tokenizer = None
22
+ _load_error: str | None = None
23
+ _loading = False
24
+ _backend: str = "none"
25
+ _active_adapter: str = "jekyll"
26
+ _base_model_id: str = "google/gemma-2-2b-it"
27
+ _last_bucket: str | None = None
28
+ _warmed_buckets: set[str] = set()
29
+
30
+
31
+ def _adapter_dirs() -> dict[AdapterName, Path]:
32
+ defaults: dict[AdapterName, Path] = {
33
+ "jekyll": ROOT / "models" / "adapters" / "jekyll-lora",
34
+ "hyde": ROOT / "models" / "adapters" / "hyde-lora",
35
+ }
36
+ if TRAIN_CONFIG.exists():
37
+ with TRAIN_CONFIG.open(encoding="utf-8") as f:
38
+ cfg = yaml.safe_load(f) or {}
39
+ adapters = cfg.get("adapters") or {}
40
+ for key in ("jekyll", "hyde"):
41
+ rel = adapters.get(key)
42
+ if rel:
43
+ defaults[key] = ROOT / rel # type: ignore[literal-required]
44
+ return defaults
45
+
46
+
47
+ def _adapter_ready(path: Path) -> bool:
48
+ return (path / "adapter_config.json").exists()
49
+
50
+
51
+ def dual_adapters_available() -> bool:
52
+ dirs = _adapter_dirs()
53
+ return _adapter_ready(dirs["jekyll"]) and _adapter_ready(dirs["hyde"])
54
+
55
+
56
+ def merged_model_available() -> bool:
57
+ return (MERGED_DIR / "config.json").exists()
58
+
59
+
60
+ def model_weights_available() -> bool:
61
+ return dual_adapters_available() or merged_model_available()
62
+
63
+
64
+ def is_loaded() -> bool:
65
+ return _model is not None and _tokenizer is not None
66
+
67
+
68
+ def is_loading() -> bool:
69
+ return _loading
70
+
71
+
72
+ def load_error() -> str | None:
73
+ return _load_error
74
+
75
+
76
+ def backend_mode() -> str:
77
+ return _backend
78
+
79
+
80
+ def active_adapter() -> AdapterName:
81
+ return _active_adapter
82
+
83
+
84
+ def resolve_adapter(persona: str | None) -> AdapterName:
85
+ focus = (persona or "balanced").lower()
86
+ if focus == "hyde":
87
+ return "hyde"
88
+ return "jekyll"
89
+
90
+
91
+ @dataclass(frozen=True)
92
+ class LocalModelInfo:
93
+ name: str
94
+ display_name: str
95
+ available: bool
96
+ fine_tuned: bool
97
+ base: str
98
+ backend: str
99
+ params_b: int | None = None
100
+ method: str = "dual-lora"
101
+ active_adapter: str = "jekyll"
102
+
103
+
104
+ def read_manifest() -> dict[str, Any]:
105
+ if MANIFEST_PATH.exists():
106
+ with MANIFEST_PATH.open(encoding="utf-8") as f:
107
+ return json.load(f)
108
+ if dual_adapters_available():
109
+ return {
110
+ "name": "jekyll-hyde",
111
+ "display_name": "Jekyll & Hyde",
112
+ "fine_tuned": True,
113
+ "base_huggingface": _base_model_id,
114
+ "base_key": "gemma2-2b",
115
+ "method": "dual-lora",
116
+ "params_b": 2,
117
+ }
118
+ if merged_model_available():
119
+ return {
120
+ "name": "jekyll-hyde",
121
+ "display_name": "Jekyll & Hyde",
122
+ "fine_tuned": True,
123
+ "base_huggingface": "google/gemma-2-2b-it",
124
+ "base_key": "gemma2-2b",
125
+ "method": "lora-merge",
126
+ "params_b": 2,
127
+ }
128
+ return {}
129
+
130
+
131
+ def get_local_model_info() -> LocalModelInfo:
132
+ manifest = read_manifest()
133
+ if not model_weights_available():
134
+ return LocalModelInfo(
135
+ name="jekyll-hyde",
136
+ display_name="Jekyll & Hyde",
137
+ available=False,
138
+ fine_tuned=False,
139
+ base="",
140
+ backend="local",
141
+ )
142
+ method = manifest.get("method", "dual-lora" if dual_adapters_available() else "lora-merge")
143
+ return LocalModelInfo(
144
+ name=manifest.get("name", "jekyll-hyde"),
145
+ display_name=manifest.get("display_name", "Jekyll & Hyde"),
146
+ available=_load_error is None or _model is not None,
147
+ fine_tuned=True,
148
+ base=manifest.get("base_huggingface", manifest.get("base", "gemma")),
149
+ backend="local",
150
+ params_b=manifest.get("params_b"),
151
+ method=method,
152
+ active_adapter=_active_adapter,
153
+ )
154
+
155
+
156
+ def normalize_messages(messages: list[dict[str, str]]) -> list[dict[str, str]]:
157
+ """Gemma 2 chat templates reject system role; fold into first user turn."""
158
+ system_parts: list[str] = []
159
+ out: list[dict[str, str]] = []
160
+ for msg in messages:
161
+ if msg["role"] == "system":
162
+ system_parts.append(msg["content"])
163
+ else:
164
+ out.append(dict(msg))
165
+ if system_parts:
166
+ prefix = "\n\n".join(system_parts)
167
+ for i, msg in enumerate(out):
168
+ if msg["role"] == "user":
169
+ out[i] = {"role": "user", "content": f"{prefix}\n\n{msg['content']}"}
170
+ break
171
+ return out
172
+
173
+
174
+ _TURN_LEAK_MARKERS = (
175
+ "<start_of_turn>",
176
+ "<end_of_turn>",
177
+ "\nuser\n",
178
+ "\nmodel\n",
179
+ "\nmodel ",
180
+ "\nassistant\n",
181
+ )
182
+
183
+
184
+ def clean_generation(text: str) -> str:
185
+ """Trim role leaks, turn markers, repeated paragraphs, and template meta from model output."""
186
+ from safety_eval.platform.output_guard import looks_like_template_leak
187
+
188
+ t = text.strip()
189
+ if not t:
190
+ return t
191
+
192
+ if looks_like_template_leak(t):
193
+ for marker in (
194
+ "Response Template",
195
+ "RESPONSE TEMPLATE",
196
+ "KEY CONCEPT",
197
+ "Example Response Template",
198
+ "SAMPLE ANSWER",
199
+ "USER QUERY:",
200
+ ):
201
+ idx = t.find(marker)
202
+ if idx >= 0:
203
+ t = t[:idx].strip()
204
+ break
205
+
206
+ lower = t.lower()
207
+ for marker in _TURN_LEAK_MARKERS:
208
+ idx = lower.find(marker.lower())
209
+ if idx > 0:
210
+ t = t[:idx].strip()
211
+ lower = t.lower()
212
+
213
+ lines: list[str] = []
214
+ for line in t.splitlines():
215
+ if line.strip().lower() in {"model", "user", "assistant"}:
216
+ break
217
+ lines.append(line)
218
+ t = "\n".join(lines).strip()
219
+
220
+ paras = [p.strip() for p in re.split(r"\n{2,}", t) if p.strip()]
221
+ deduped: list[str] = []
222
+ for p in paras:
223
+ if not deduped or p != deduped[-1]:
224
+ deduped.append(p)
225
+ return "\n\n".join(deduped).strip()
226
+
227
+
228
+ def _load_dual_adapters() -> tuple[Any, Any]:
229
+ global _backend, _active_adapter, _base_model_id
230
+
231
+ import sys
232
+
233
+ if str(ROOT) not in sys.path:
234
+ sys.path.insert(0, str(ROOT))
235
+ from training.bootstrap_adapters import bootstrap_dual_adapters
236
+
237
+ bootstrap_dual_adapters()
238
+ if not dual_adapters_available():
239
+ raise RuntimeError("Dual LoRA adapters missing. Run training/train_lora.py --persona both")
240
+
241
+ import os
242
+
243
+ os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
244
+ os.environ.setdefault("TRANSFORMERS_NO_ADVISORY_WARNINGS", "1")
245
+
246
+ import torch
247
+ from peft import PeftModel
248
+ from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
249
+
250
+ manifest = read_manifest()
251
+ model_id = manifest.get("base_huggingface", _base_model_id)
252
+ _base_model_id = model_id
253
+ dirs = _adapter_dirs()
254
+
255
+ tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
256
+ if tokenizer.pad_token is None:
257
+ tokenizer.pad_token = tokenizer.eos_token
258
+
259
+ if torch.cuda.is_available():
260
+ quant = BitsAndBytesConfig(
261
+ load_in_4bit=True,
262
+ bnb_4bit_quant_type="nf4",
263
+ bnb_4bit_compute_dtype=torch.float16,
264
+ bnb_4bit_use_double_quant=True,
265
+ )
266
+ base = AutoModelForCausalLM.from_pretrained(
267
+ model_id,
268
+ trust_remote_code=True,
269
+ quantization_config=quant,
270
+ device_map="auto",
271
+ )
272
+ else:
273
+ base = AutoModelForCausalLM.from_pretrained(
274
+ model_id,
275
+ trust_remote_code=True,
276
+ torch_dtype=torch.float32,
277
+ device_map="cpu",
278
+ )
279
+
280
+ model = PeftModel.from_pretrained(base, str(dirs["jekyll"]), adapter_name="jekyll")
281
+ model.load_adapter(str(dirs["hyde"]), adapter_name="hyde")
282
+ model.set_adapter("jekyll")
283
+ model.eval()
284
+ _backend = "dual-lora"
285
+ _active_adapter = "jekyll"
286
+ return model, tokenizer
287
+
288
+
289
+ def _load_merged() -> tuple[Any, Any]:
290
+ global _backend
291
+
292
+ import os
293
+
294
+ os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
295
+ os.environ.setdefault("TRANSFORMERS_NO_ADVISORY_WARNINGS", "1")
296
+
297
+ import torch
298
+ from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
299
+
300
+ tokenizer = AutoTokenizer.from_pretrained(MERGED_DIR, trust_remote_code=True)
301
+ if tokenizer.pad_token is None:
302
+ tokenizer.pad_token = tokenizer.eos_token
303
+
304
+ if torch.cuda.is_available():
305
+ quant = BitsAndBytesConfig(
306
+ load_in_4bit=True,
307
+ bnb_4bit_quant_type="nf4",
308
+ bnb_4bit_compute_dtype=torch.float16,
309
+ bnb_4bit_use_double_quant=True,
310
+ )
311
+ model = AutoModelForCausalLM.from_pretrained(
312
+ MERGED_DIR,
313
+ trust_remote_code=True,
314
+ quantization_config=quant,
315
+ device_map="auto",
316
+ )
317
+ else:
318
+ model = AutoModelForCausalLM.from_pretrained(
319
+ MERGED_DIR,
320
+ trust_remote_code=True,
321
+ torch_dtype=torch.float32,
322
+ device_map="cpu",
323
+ )
324
+
325
+ model.eval()
326
+ _backend = "merged"
327
+ return model, tokenizer
328
+
329
+
330
+ def set_lora_mix(jekyll_w: float, hyde_w: float) -> None:
331
+ """Blend jekyll + hyde LoRA adapters using pre-warmed MoE bucket pool."""
332
+ global _active_adapter, _last_bucket
333
+
334
+ if _model is None or _backend != "dual-lora":
335
+ _set_active_adapter(resolve_adapter("hyde" if hyde_w > jekyll_w else "jekyll"))
336
+ return
337
+
338
+ from safety_eval.platform.lora_mix_cache import MOE_BUCKETS, record_mix_usage, snap_to_bucket
339
+
340
+ snap = snap_to_bucket(jekyll_w, hyde_w)
341
+ record_mix_usage(snap)
342
+
343
+ if _last_bucket == snap.bucket_id:
344
+ return
345
+
346
+ if snap.adapter_name in ("jekyll", "hyde"):
347
+ _model.set_adapter(snap.adapter_name)
348
+ _active_adapter = snap.adapter_name # type: ignore[assignment]
349
+ _last_bucket = snap.bucket_id
350
+ return
351
+
352
+ try:
353
+ if not _adapter_exists(snap.adapter_name):
354
+ _model.add_weighted_adapter(
355
+ adapters=["jekyll", "hyde"],
356
+ weights=[snap.jekyll, snap.hyde],
357
+ adapter_name=snap.adapter_name,
358
+ combination_type="linear",
359
+ )
360
+ _warmed_buckets.add(snap.adapter_name)
361
+ _model.set_adapter(snap.adapter_name)
362
+ _active_adapter = snap.adapter_name # type: ignore[assignment]
363
+ _last_bucket = snap.bucket_id
364
+ except Exception:
365
+ _set_active_adapter("jekyll" if snap.jekyll >= snap.hyde else "hyde")
366
+ _last_bucket = snap.bucket_id
367
+
368
+
369
+ def _adapter_exists(name: str) -> bool:
370
+ if _model is None:
371
+ return False
372
+ return name in getattr(_model, "peft_config", {})
373
+
374
+
375
+ def prewarm_moe_buckets() -> int:
376
+ """Pre-create all five MoE bucket adapters to avoid per-request overhead."""
377
+ if _model is None or _backend != "dual-lora":
378
+ return 0
379
+ warmed = 0
380
+ for name, jw, hw in MOE_BUCKETS:
381
+ if _adapter_exists(name):
382
+ _warmed_buckets.add(name)
383
+ continue
384
+ try:
385
+ _model.add_weighted_adapter(
386
+ adapters=["jekyll", "hyde"],
387
+ weights=[jw, hw],
388
+ adapter_name=name,
389
+ combination_type="linear",
390
+ )
391
+ _warmed_buckets.add(name)
392
+ warmed += 1
393
+ except Exception:
394
+ continue
395
+ _model.set_adapter("jekyll")
396
+ _active_adapter = "jekyll"
397
+ return warmed
398
+
399
+
400
+ def _set_active_adapter(name: AdapterName) -> None:
401
+ global _active_adapter
402
+ if _model is None or _backend != "dual-lora":
403
+ return
404
+ _model.set_adapter(name)
405
+ _active_adapter = name
406
+
407
+
408
+ def _ensure_loaded() -> tuple[Any, Any]:
409
+ global _model, _tokenizer, _load_error
410
+
411
+ if _model is not None and _tokenizer is not None:
412
+ return _model, _tokenizer
413
+
414
+ if not model_weights_available():
415
+ raise RuntimeError(
416
+ "Fine-tuned model not found. Run training/train_lora.py --persona both then merge."
417
+ )
418
+
419
+ try:
420
+ import torch # noqa: F401
421
+ from transformers import AutoModelForCausalLM # noqa: F401
422
+ except ImportError as exc:
423
+ _load_error = "Install training env: pip install -e '.[train]' (or use .venv-train)"
424
+ raise RuntimeError(_load_error) from exc
425
+
426
+ try:
427
+ if dual_adapters_available():
428
+ _model, _tokenizer = _load_dual_adapters()
429
+ else:
430
+ _model, _tokenizer = _load_merged()
431
+ _load_error = None
432
+ return _model, _tokenizer
433
+ except Exception as exc:
434
+ _load_error = str(exc)
435
+ raise RuntimeError(f"Failed to load fine-tuned model: {exc}") from exc
436
+
437
+
438
+ def chat(
439
+ messages: list[dict[str, str]],
440
+ *,
441
+ temperature: float = 0.7,
442
+ max_new_tokens: int = 384,
443
+ adapter: str | None = None,
444
+ lora_mix: tuple[float, float] | None = None,
445
+ grammar: str | None = None,
446
+ ) -> str:
447
+ import torch
448
+
449
+ from safety_eval.platform.decoding_entropy import apply_to_generation_kwargs, decoding_for_lora_mix
450
+
451
+ model, tokenizer = _ensure_loaded()
452
+ mix_j, mix_h = 1.0, 0.0
453
+ if lora_mix is not None:
454
+ set_lora_mix(lora_mix[0], lora_mix[1])
455
+ mix_j, mix_h = lora_mix[0], lora_mix[1]
456
+ elif adapter:
457
+ _set_active_adapter(resolve_adapter(adapter))
458
+ mix_j = 1.0 if resolve_adapter(adapter) == "jekyll" else 0.0
459
+ mix_h = 1.0 - mix_j
460
+
461
+ decode = decoding_for_lora_mix(mix_j, mix_h, base_temperature=temperature)
462
+
463
+ norm = normalize_messages(messages)
464
+ prompt = tokenizer.apply_chat_template(norm, tokenize=False, add_generation_prompt=True)
465
+ inputs = tokenizer(prompt, return_tensors="pt")
466
+ device = next(model.parameters()).device
467
+ inputs = {k: v.to(device) for k, v in inputs.items()}
468
+
469
+ gen_kwargs: dict[str, Any] = {
470
+ "max_new_tokens": max_new_tokens,
471
+ "pad_token_id": tokenizer.pad_token_id or tokenizer.eos_token_id,
472
+ "eos_token_id": tokenizer.eos_token_id,
473
+ "repetition_penalty": 1.12,
474
+ "no_repeat_ngram_size": 4,
475
+ }
476
+ gen_kwargs = apply_to_generation_kwargs(decode, gen_kwargs)
477
+
478
+ if grammar == "mcp_tool_json":
479
+ from safety_eval.platform.grammar_constraint import build_mcp_tool_prefix_fn
480
+
481
+ gen_kwargs["prefix_allowed_tokens_fn"] = build_mcp_tool_prefix_fn(tokenizer)
482
+
483
+ with torch.no_grad():
484
+ output = model.generate(**inputs, **gen_kwargs)
485
+
486
+ new_tokens = output[0][inputs["input_ids"].shape[1] :]
487
+ return clean_generation(tokenizer.decode(new_tokens, skip_special_tokens=True))
488
+
489
+
490
+ def reload_model() -> LocalModelInfo:
491
+ """Unload and reload weights after incremental training."""
492
+ global _model, _tokenizer, _load_error, _loading, _backend, _active_adapter, _last_bucket, _warmed_buckets
493
+ _model = None
494
+ _tokenizer = None
495
+ _load_error = None
496
+ _loading = False
497
+ _backend = "none"
498
+ _active_adapter = "jekyll"
499
+ _last_bucket = None
500
+ _warmed_buckets = set()
501
+ return preload()
502
+
503
+
504
+ def preload() -> LocalModelInfo:
505
+ """Warm up GPU weights (call from background thread)."""
506
+ global _loading, _load_error
507
+ if not model_weights_available():
508
+ return get_local_model_info()
509
+ if is_loaded():
510
+ return get_local_model_info()
511
+ _loading = True
512
+ _load_error = None
513
+ try:
514
+ _ensure_loaded()
515
+ if _backend == "dual-lora":
516
+ import threading
517
+
518
+ threading.Thread(target=prewarm_moe_buckets, daemon=True).start()
519
+ except Exception as exc:
520
+ _load_error = str(exc)
521
+ raise
522
+ finally:
523
+ _loading = False
524
+ return get_local_model_info()