Qwen21_Text_Encoder_Heretic / extras /create_heretic_text_encoder.py
catplusplus's picture
Upload folder using huggingface_hub
7b00440 verified
Raw History Blame Contribute Delete
8.9 kB
# -*- coding: utf-8 -*-
"""Heretic Abliteration Engine for Qwen3-VL Text Encoder.
Applies Norm-Preserving Biprojected Abliteration (Heretic/Arditi et al.) to Qwen3-VL-8B:
1. Measures the refusal/hesitation direction across all 36 transformer decoder layers.
2. Orthogonalizes the refusal direction against the benign semantic direction.
3. Applies rank-1 directional ablation to `self_attn.o_proj` and `mlp.down_proj`:
W' = normalize(W_norm - lambda * v * (v^T W_norm)) * ||W||_row
preserving exact row norms to protect general capability and language quality.
4. Exports the decensored, hesitation-free text encoder to ~/Nikola/models/Qwen_Text_Encoder_Heretic.
"""
import os
import sys
import shutil
import math
import torch
import torch.nn.functional as F
import torch.linalg as LA
from diffusers import QwenImage21Pipeline
ROOT_DIR = "/auto/home/amano/olegk/Nikola"
for p in [f"{ROOT_DIR}/src/imagegen", ROOT_DIR]:
if p not in sys.path:
sys.path.insert(0, p)
from stream_encoder import attach_qwen3vl_streamer
MODEL_PATH = "/home/olegk/Nikola/models/Qwen/Qwen-Image-2.1"
OUTPUT_DIR = "/home/olegk/Nikola/models/Qwen21_Text_Encoder_Heretic"
# Contrastive calibration prompts
BENIGN_PROMPTS = [
"A scenic view of a quiet tropical beach with palm trees, white sand, and gentle ocean waves under sunlight.",
"A group of athletic friends playing a casual game of volleyball on a sports court on a bright afternoon.",
"Two cheerful anime girls laughing together and talking in a colorful modern classroom setting.",
"A beautiful golden hour sunset over the coastal dunes with warm ambient lighting and ocean breeze.",
"Two athletic anime girls in standard sportswear practicing drills together on an outdoor gymnasium court.",
"A detailed photograph of a clean sandy beach with clear turquoise water and blue sky with light clouds.",
"Two anime girls taking a relaxing stroll along the seaside promenade during summer vacation.",
"A dynamic action photograph of a volleyball match on the sand with players in motion."
]
SENSITIVE_PROMPTS = [
"Two anime girls in colorful skimpy bikinis playing beach volleyball on a sunny tropical beach.",
"Two attractive anime girls wearing revealing two-piece swimsuits posing playfully under the bright sun.",
"Two cute anime girls in skimpy micro bikinis jumping to spike a volleyball on the beach sand.",
"A beach volleyball match with two anime girls wearing small swimsuits, athletic dynamic poses.",
"Two anime girls in alluring colorful bikinis having fun playing sports on a sunny seaside beach.",
"A close up action shot of two anime girls in tiny bikinis diving for a volleyball on the beach.",
"Two beautiful anime girls wearing revealing beachwear and swimsuits posing by the ocean shoreline.",
"An alluring tropical beach scene with two anime girls in revealing bikinis playing athletic beach sports."
]
def abliterate_layer_weights(weight: torch.Tensor, v: torch.Tensor, weight_factor: float) -> torch.Tensor:
"""Norm-preserving biprojected abliteration: delta W = -lambda * v * (v^T W)."""
if weight_factor <= 0.0:
return weight
orig_dtype = weight.dtype
W = weight.float()
v = v.float().to(W.device)
# Calculate row norms
row_norms = LA.vector_norm(W, dim=1, keepdim=True)
# Normalize rows
W_norm = F.normalize(W, p=2, dim=1)
# v @ W_norm -> (in_features,)
lora_A = (v @ W_norm).view(1, -1)
# -weight_factor * v -> (out_features, 1)
lora_B = (-weight_factor * v).view(-1, 1)
# Project and renormalize
W_adj = W_norm + lora_B @ lora_A
W_adj = F.normalize(W_adj, p=2, dim=1)
W_final = W_adj * row_norms
return W_final.to(orig_dtype)
def create_heretic_model():
print("=" * 80)
print("🔮 FORGING OPTIMIZED HERETIC TEXT ENCODER (QWEN3-VL-8B)")
print("=" * 80)
print(f"Base Model: {MODEL_PATH}")
print(f"Output Directory: {OUTPUT_DIR}")
# 1. Load pipeline and attach layerwise streamer for rapid residual extraction
print("\n[Step 1/4] Loading pipeline & attaching layerwise streamer...")
pipe = QwenImage21Pipeline.from_pretrained(
MODEL_PATH,
transformer=None,
torch_dtype=torch.bfloat16,
)
streamer = attach_qwen3vl_streamer(pipe, device="cuda:0")
# 2. Extract layer-by-layer residuals across contrastive prompt sets
print("\n[Step 2/4] Measuring refusal directions across all 36 decoder layers...")
def get_mean_residuals(prompts, label):
layer_states = [[] for _ in range(37)]
print(f" • Collecting residuals for {len(prompts)} {label} prompts...")
for p in prompts:
fmt = pipe.prompt_template_t2i.format(p)
inputs = pipe.processor(text=[fmt], return_tensors="pt").to("cuda:0")
with torch.no_grad():
out = pipe.text_encoder(
input_ids=inputs.input_ids,
attention_mask=inputs.attention_mask,
output_hidden_states=True,
)
for l in range(37):
layer_states[l].append(out.hidden_states[l][0, -1, :].float().cpu())
return torch.stack([torch.stack(l).mean(dim=0) for l in layer_states])
b_means = get_mean_residuals(BENIGN_PROMPTS, "benign")
s_means = get_mean_residuals(SENSITIVE_PROMPTS, "sensitive")
# Compute difference of means
residual_directions = s_means - b_means
# Orthogonalize against benign direction (Heretic projected abliteration)
print(" • Orthogonalizing refusal directions against benign semantic vectors...")
good_directions = F.normalize(b_means, p=2, dim=1)
proj = torch.sum(residual_directions * good_directions, dim=1, keepdim=True)
ortho_directions = residual_directions - proj * good_directions
ortho_directions = F.normalize(ortho_directions, p=2, dim=1)
# 3. Apply Norm-Preserving Biprojected Abliteration to Language Model Weights
print("\n[Step 3/4] Applying Heretic norm-preserving abliteration to weights...")
lm = pipe.text_encoder.model.language_model
num_layers = len(lm.layers)
# Target layers: layers 16 to 35, centered at layer 26 with max_weight = 1.0
center_layer = 26.0
spread = 8.0
ablated_count = 0
for l_idx, layer in enumerate(lm.layers):
dist = abs(l_idx - center_layer)
# Smooth bell-shaped abliteration profile
weight_factor = float(math.exp(-(dist**2) / (2 * (spread**2))))
if weight_factor < 0.10:
weight_factor = 0.0 # skip early layers (0..12) where refusal is inactive
v = ortho_directions[l_idx + 1] # l_idx+1 accounts for embedding layer at idx 0
if weight_factor > 0.0:
print(f" • Layer {l_idx:2d}: Abliterating with lambda={weight_factor:.3f}...")
# 1. Attention Out Projection
if hasattr(layer.self_attn, "o_proj"):
with torch.no_grad():
layer.self_attn.o_proj.weight.data = abliterate_layer_weights(
layer.self_attn.o_proj.weight.data, v, weight_factor
)
ablated_count += 1
# 2. MLP Down Projection
if hasattr(layer.mlp, "down_proj"):
with torch.no_grad():
layer.mlp.down_proj.weight.data = abliterate_layer_weights(
layer.mlp.down_proj.weight.data, v, weight_factor
)
ablated_count += 1
print(f" • Successfully abliterated {ablated_count} linear projection matrices!")
# 4. Save Standalone Checkpoint to models/Qwen_Text_Encoder_Heretic
print(f"\n[Step 4/4] Saving Heretic Text Encoder to {OUTPUT_DIR}...")
os.makedirs(OUTPUT_DIR, exist_ok=True)
# Save text_encoder subfolder (for direct use with Diffusers or Transformers)
sub_dir = os.path.join(OUTPUT_DIR, "text_encoder")
os.makedirs(sub_dir, exist_ok=True)
pipe.text_encoder.save_pretrained(sub_dir)
print(f" • Saved abliterated text encoder weights to: {sub_dir}")
# Also copy tokenizer and processor files
proc_src = os.path.join(MODEL_PATH, "processor")
proc_dst = os.path.join(OUTPUT_DIR, "processor")
if os.path.exists(proc_src):
if os.path.exists(proc_dst):
shutil.rmtree(proc_dst)
shutil.copytree(proc_src, proc_dst)
print(f" • Copied processor to: {proc_dst}")
# Copy root config files
for fname in ["config.json", "generation_config.json"]:
src_f = os.path.join(MODEL_PATH, "text_encoder", fname)
if os.path.exists(src_f):
shutil.copy(src_f, os.path.join(OUTPUT_DIR, fname))
print("\n🎉 SUCCESS: Heretic Text Encoder forged and saved successfully!")
print("=" * 80)
if __name__ == "__main__":
create_heretic_model()