File size: 8,895 Bytes
7b00440
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
# -*- 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()