pliny-the-prompter commited on
Commit
4f13809
·
verified ·
1 Parent(s): 0fdea90

Upload 133 files

Browse files
Files changed (2) hide show
  1. app.py +4 -3
  2. obliteratus/abliterate.py +9 -3
app.py CHANGED
@@ -2661,7 +2661,8 @@ def chat_respond(message: str, history: list[dict], system_prompt: str,
2661
  text = "\n".join(f"{m['role']}: {m['content']}" for m in messages) + "\nassistant:"
2662
 
2663
  inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=context_length)
2664
- inputs = {k: v.to(model.device) for k, v in inputs.items()}
 
2665
 
2666
  # Streaming generation — repetition_penalty (user-controllable, default 1.0)
2667
  # can break degenerate refusal loops if increased.
@@ -3156,7 +3157,7 @@ def ab_chat_respond(message: str, history_left: list[dict], history_right: list[
3156
  # --- Generate from abliterated model (streaming) ---
3157
  stream_timeout = max(120, 120 + int(max_tokens * 0.1))
3158
  streamer_abl = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=stream_timeout)
3159
- inputs_abl = {k: v.to(abliterated_model.device) for k, v in inputs.items()}
3160
  gen_kwargs_abl = {**inputs_abl, **gen_kwargs_base, "streamer": streamer_abl}
3161
 
3162
  gen_error_abl = [None]
@@ -3217,7 +3218,7 @@ def ab_chat_respond(message: str, history_left: list[dict], history_right: list[
3217
  )
3218
 
3219
  streamer_orig = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=stream_timeout)
3220
- inputs_orig = {k: v.to(original_model.device) for k, v in inputs.items()}
3221
  gen_kwargs_orig = {**inputs_orig, **gen_kwargs_base, "streamer": streamer_orig}
3222
 
3223
  gen_error_orig = [None]
 
2661
  text = "\n".join(f"{m['role']}: {m['content']}" for m in messages) + "\nassistant:"
2662
 
2663
  inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=context_length)
2664
+ _model_device = next(model.parameters()).device
2665
+ inputs = {k: v.to(_model_device) for k, v in inputs.items()}
2666
 
2667
  # Streaming generation — repetition_penalty (user-controllable, default 1.0)
2668
  # can break degenerate refusal loops if increased.
 
3157
  # --- Generate from abliterated model (streaming) ---
3158
  stream_timeout = max(120, 120 + int(max_tokens * 0.1))
3159
  streamer_abl = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=stream_timeout)
3160
+ inputs_abl = {k: v.to(next(abliterated_model.parameters()).device) for k, v in inputs.items()}
3161
  gen_kwargs_abl = {**inputs_abl, **gen_kwargs_base, "streamer": streamer_abl}
3162
 
3163
  gen_error_abl = [None]
 
3218
  )
3219
 
3220
  streamer_orig = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=stream_timeout)
3221
+ inputs_orig = {k: v.to(next(original_model.parameters()).device) for k, v in inputs.items()}
3222
  gen_kwargs_orig = {**inputs_orig, **gen_kwargs_base, "streamer": streamer_orig}
3223
 
3224
  gen_error_orig = [None]
obliteratus/abliterate.py CHANGED
@@ -1052,7 +1052,8 @@ class AbliterationPipeline:
1052
  else:
1053
  # Layer produced no activations (hook failure or skipped layer)
1054
  empty_layers.append(idx)
1055
- hidden = self._harmful_acts[0][0].shape[-1] if self._harmful_acts.get(0) else 768
 
1056
  self._harmful_means[idx] = torch.zeros(1, hidden)
1057
  self._harmless_means[idx] = torch.zeros(1, hidden)
1058
  if empty_layers:
@@ -1072,7 +1073,8 @@ class AbliterationPipeline:
1072
  if self._jailbreak_acts.get(idx):
1073
  self._jailbreak_means[idx] = torch.stack(self._jailbreak_acts[idx]).mean(dim=0)
1074
  else:
1075
- hidden = self._harmful_acts[0][0].shape[-1] if self._harmful_acts.get(0) else 768
 
1076
  self._jailbreak_means[idx] = torch.zeros(1, hidden)
1077
  self.log(" Jailbreak activations collected for three-way contrastive analysis")
1078
 
@@ -1138,15 +1140,19 @@ class AbliterationPipeline:
1138
  pass # Fall through to per-prompt with error handling
1139
 
1140
  wrapped = []
 
1141
  for i, conv in enumerate(all_conversations):
1142
  try:
1143
  text = tokenizer.apply_chat_template(
1144
  conv, tokenize=False, add_generation_prompt=True
1145
  )
1146
  wrapped.append(text)
 
1147
  except Exception:
1148
  wrapped.append(prompts[i]) # fallback to raw if individual prompt fails
1149
- self.log(f" chat template {n}/{n}")
 
 
1150
  return wrapped
1151
 
1152
  def _apply_spectral_cascade_weights(self):
 
1052
  else:
1053
  # Layer produced no activations (hook failure or skipped layer)
1054
  empty_layers.append(idx)
1055
+ _acts_0 = self._harmful_acts.get(0)
1056
+ hidden = _acts_0[0].shape[-1] if _acts_0 and len(_acts_0) > 0 else 768
1057
  self._harmful_means[idx] = torch.zeros(1, hidden)
1058
  self._harmless_means[idx] = torch.zeros(1, hidden)
1059
  if empty_layers:
 
1073
  if self._jailbreak_acts.get(idx):
1074
  self._jailbreak_means[idx] = torch.stack(self._jailbreak_acts[idx]).mean(dim=0)
1075
  else:
1076
+ _acts_0 = self._harmful_acts.get(0)
1077
+ hidden = _acts_0[0].shape[-1] if _acts_0 and len(_acts_0) > 0 else 768
1078
  self._jailbreak_means[idx] = torch.zeros(1, hidden)
1079
  self.log(" Jailbreak activations collected for three-way contrastive analysis")
1080
 
 
1140
  pass # Fall through to per-prompt with error handling
1141
 
1142
  wrapped = []
1143
+ n_wrapped = 0
1144
  for i, conv in enumerate(all_conversations):
1145
  try:
1146
  text = tokenizer.apply_chat_template(
1147
  conv, tokenize=False, add_generation_prompt=True
1148
  )
1149
  wrapped.append(text)
1150
+ n_wrapped += 1
1151
  except Exception:
1152
  wrapped.append(prompts[i]) # fallback to raw if individual prompt fails
1153
+ self.log(f" chat template {n_wrapped}/{n}")
1154
+ if n_wrapped < n:
1155
+ self.log(f" WARNING: {n - n_wrapped} prompts fell back to raw text")
1156
  return wrapped
1157
 
1158
  def _apply_spectral_cascade_weights(self):