Spaces:
Sleeping
Sleeping
File size: 7,309 Bytes
4d362a0 | 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 | """fase g — gauge-flip valor-salida demostrado, sin reentrenar.
aplica una transformación de gauge (en convención del paper v5:
w_v <- w_v r, w_o <- r^{-1} w_o, r en gl(d_h)) que deja la salida intacta
a precisión de máquina, y muestra que el ranking de redundancia por
v1(w_o) ---primer vector singular derecho de w_o, en r^768--- cambia
mientras la firma del circuito ov (w_v w_o) queda estable. convierte la
crítica de la premisa en una vulnerabilidad exhibida: leer la dirección
de w_o es leer ruido de gauge.
"""
import torch
from src.firma_funcional import w_o_por_cabeza, w_v_por_cabeza
@torch.no_grad()
def aplica_gauge_ov(
modelo,
capa: int,
n_cabezas: int,
dim_cabeza: int,
semilla: int = 0,
escala_id: float | None = None,
) -> float:
"""aplica un gauge valor-salida por cabeza, in situ.
sustituye w_v <- r^t w_v, b_v <- r^t b_v y w_o <- w_o r^{-t} con r en
gl(d_h) bien condicionada; deja la salida intacta a precisión de
máquina y denota la libertad de gauge del sector valor-salida. el
sesgo de valor entra en la transformación porque v = x w_v^t + b_v y
la compensación de w_o exige que b_v gire igual que la weight; en una
columna con qkv_bias omitirlo rompe la invariancia. clonar el modelo
antes si se quiere conservar el original.
args:
modelo: el vit, modificado in situ.
capa: índice de la capa intervenida.
n_cabezas: cabezas h.
dim_cabeza: d_h.
semilla: semilla de la r aleatoria, para reproducir.
escala_id: coeficiente de la identidad en r = randn + escala_id*i;
menor valor aleja r de un múltiplo escalar de la identidad
---gauge más fuerte, v1(w_o) deriva más, el circuito ov no---.
por defecto sqrt(d_h), el calibre de referencia.
returns:
desviación media de r respecto a su mejor múltiplo escalar de la
identidad, sobre las cabezas; proxy de la fuerza del gauge.
"""
if escala_id is None:
escala_id = dim_cabeza ** 0.5
g = torch.Generator(device="cpu").manual_seed(semilla)
attn = modelo.blocks[capa].attn
w_qkv = attn.qkv.weight # [3d, d]
w_proj = attn.proj.weight # [d, d]
base = 2 * w_qkv.shape[1] # inicio de v
# r e inversa en float64: con r fuerte la inv en fp32 amplifica el
# error por el número de condición y rompe la invariancia; en fp64 el
# gauge es preciso sea cual sea la fuerza, y solo se pierde el último
# casteo al dtype del peso.
ident = torch.eye(dim_cabeza, device=w_qkv.device, dtype=torch.float64)
desvs = []
for h in range(n_cabezas):
rt = torch.randn(dim_cabeza, dim_cabeza, generator=g,
dtype=torch.float64).to(w_qkv.device)
r = rt + escala_id * ident
# desviación de r respecto a su mejor múltiplo escalar de la id
s = r.diagonal().mean()
desvs.append(float((r - s * ident).norm()
/ (s.abs() * dim_cabeza ** 0.5 + 1e-8)))
r_inv_t = torch.linalg.inv(r).t()
fil = slice(base + h * dim_cabeza, base + (h + 1) * dim_cabeza)
col = slice(h * dim_cabeza, (h + 1) * dim_cabeza)
w_qkv[fil, :] = (
r.t() @ w_qkv[fil, :].double()).to(w_qkv.dtype) # w_v
if attn.qkv.bias is not None:
attn.qkv.bias[fil] = (
r.t() @ attn.qkv.bias[fil].double()).to(w_qkv.dtype) # b_v
w_proj[:, col] = (
w_proj[:, col].double() @ r_inv_t).to(w_proj.dtype) # w_o
return float(sum(desvs) / len(desvs))
@torch.no_grad()
def cos_pares_v1_wo(
modelo,
capa: int,
n_cabezas: int,
dim_cabeza: int,
) -> torch.Tensor:
"""|cos| entre las direcciones v1(w_o) de las cabezas de una capa.
v1(w_o) es el primer vector singular derecho de w_o (en r^768), la
cantidad que la poda o interpretación por dirección de salida lee; el
experimento muestra que cambia bajo el gauge.
args:
modelo: el vit.
capa: índice de la capa.
n_cabezas: cabezas h.
dim_cabeza: d_h.
returns:
tensor [h, h] con |cos| entre direcciones dominantes de w_o.
"""
w_o = w_o_por_cabeza(modelo, capa, n_cabezas, dim_cabeza) # [h,dh,d]
u = torch.stack([
torch.linalg.svd(w_o[h], full_matrices=False).Vh[0]
for h in range(n_cabezas)]) # [h, d]
return (u @ u.t()).abs()
@torch.no_grad()
def cos_pares_circuito_ov(
modelo,
capa: int,
n_cabezas: int,
dim_cabeza: int,
) -> torch.Tensor:
"""|cos| entre las firmas invariantes del circuito ov por cabeza.
contrasta con cos_pares_v1_wo: no cambia bajo el gauge. el
diagnóstico usa svd exacta del circuito ov (w_v w_o) ---no la
iteración de potencia del regularizador---, de modo que la
invariancia se ve a precisión de máquina y no arrastra el ruido de
init de la potencia.
args:
modelo: el vit.
capa: índice de la capa.
n_cabezas: cabezas h.
dim_cabeza: d_h.
returns:
tensor [h, h] con |cos| entre firmas del circuito ov.
"""
w_v = w_v_por_cabeza(modelo, capa, n_cabezas, dim_cabeza) # [h,dh,d]
w_o = w_o_por_cabeza(modelo, capa, n_cabezas, dim_cabeza) # [h,dh,d]
# circuito ov por cabeza m = w_v w_o (rango <= dh, invariante de
# gauge); su primer vector singular es la firma invariante. el einsum
# contrae el eje de valor d_h compartido (w_v es [dh,d]=W_v^T)
circ = torch.einsum("hed,hef->hdf", w_v, w_o) # [h, d, d]
r = torch.stack([
torch.linalg.svd(circ[h], full_matrices=False).U[:, 0]
for h in range(n_cabezas)]) # [h, d]
return (r @ r.t()).abs()
@torch.no_grad()
def aplica_gauge_ortogonal(
modelo, capa: int, n_cabezas: int, dim_cabeza: int, semilla: int
) -> None:
"""aplica un gauge ortogonal por cabeza, in situ.
espeja `aplica_gauge_ov` salvo que r se muestrea en o(d_h) en vez
de en gl(d_h). sirve de assert de c1: bajo este gauge no debe
moverse ni v1(w_o) ni el circuito.
Args:
modelo: el vit.
capa: índice de la capa.
n_cabezas: cabezas h.
dim_cabeza: d_h.
semilla: semilla del generador de r.
"""
g = torch.Generator(device="cpu").manual_seed(semilla)
attn = modelo.blocks[capa].attn
w_qkv, w_proj = attn.qkv.weight, attn.proj.weight
base = 2 * w_qkv.shape[1]
for h in range(n_cabezas):
m = torch.randn(dim_cabeza, dim_cabeza, generator=g,
dtype=torch.float64)
r = torch.linalg.qr(m)[0] # ortogonal exacta
r_inv_t = r # (r^-1)^t = r
fil = slice(base + h * dim_cabeza,
base + (h + 1) * dim_cabeza)
col = slice(h * dim_cabeza, (h + 1) * dim_cabeza)
w_qkv[fil, :] = (r.t() @ w_qkv[fil, :].double()).to(w_qkv.dtype)
if attn.qkv.bias is not None:
attn.qkv.bias[fil] = (
r.t() @ attn.qkv.bias[fil].double()).to(w_qkv.dtype)
w_proj[:, col] = (
w_proj[:, col].double() @ r_inv_t).to(w_proj.dtype)
|