File size: 13,000 Bytes
6ee9587 | 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 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 | import torch
import typing
from quantize import utils
import transformers
import tqdm, math
from models.hadamard_utils import random_hadamard_matrix
from dataclasses import dataclass
OPT_MODEL = transformers.models.opt.modeling_opt.OPTForCausalLM
OPT_LAYER = transformers.models.opt.modeling_opt.OPTDecoderLayer
LLAMA_MODEL = transformers.models.llama.modeling_llama.LlamaForCausalLM
LLAMA_LAYER = transformers.models.llama.modeling_llama.LlamaDecoderLayer
QWEN_MODEL = transformers.models.qwen2.modeling_qwen2.Qwen2ForCausalLM
QWEN_LATER = transformers.models.qwen2.modeling_qwen2.Qwen2DecoderLayer
DEV = torch.device('cuda:0') if torch.cuda.is_available() else torch.device('cpu')
class RMSN(torch.nn.Module):
"""
This class implements the Root Mean Square Normalization (RMSN) layer.
We use the implementation from LLAMARMSNorm here:
https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L75
"""
def __init__(self, mean_dim: int,ori_norm, eps=1e-5):
super().__init__()
self.eps = eps
self.mean_dim = mean_dim
if ori_norm:
self.weight = torch.nn.Parameter(torch.ones_like(ori_norm.weight))
if hasattr(ori_norm,"bias") and ori_norm.bias is not None:
self.bias = torch.nn.Parameter(torch.zeros_like(ori_norm.bias))
else:
self.weight = torch.nn.Parameter(torch.ones(1))
self.bias = None
def forward(self, x: torch.Tensor) -> torch.Tensor:
input_dtype = x.dtype
if x.dtype == torch.float16:
x = x.to(torch.float32)
variance = x.pow(2).mean(-1, keepdim=True)
x = x * torch.rsqrt(variance + self.eps)
return (self.weight * x).to(input_dtype)
def replace_modules(
root: torch.nn.Module,
type_to_replace,
new_module_factory,
replace_layers: bool,
) -> None:
"""Replace modules of given type using the supplied module factory.
Perform a depth-first search of a module hierarchy starting at root
and replace all instances of type_to_replace with modules created by
new_module_factory. Children of replaced modules are not processed.
Args:
root: the root of the module hierarchy where modules should be replaced
type_to_replace: a type instances of which will be replaced
new_module_factory: a function that given a module that should be replaced
produces a module to replace it with.
"""
for name, module in root.named_children():
new_module = None
if isinstance(module, type_to_replace):
if replace_layers: # layernorm_fusion.replace_layers case where transformer layers are replaced
new_module = new_module_factory(module, int(name))
else: # layernorm_fusion.fuse_modules case where layernorms are fused
new_module = new_module_factory(module)
elif len(list(module.children())) > 0:
replace_modules(module, type_to_replace, new_module_factory, replace_layers)
if new_module is not None:
setattr(root, name, new_module)
def get_model_type(model):
if isinstance(model, OPT_MODEL):
model_type = OPT_MODEL
elif isinstance(model, LLAMA_MODEL):
model_type = LLAMA_MODEL
elif isinstance(model, QWEN_MODEL):
model_type = QWEN_MODEL
else:
raise ValueError(f'Unknown model type {model}')
return model_type
def model_type_extractor(model):
if isinstance(model, LLAMA_MODEL):
return LLAMA_MODEL
elif isinstance(model, QWEN_MODEL):
return QWEN_MODEL
elif isinstance(model, OPT_MODEL):
return OPT_MODEL
else:
raise ValueError(f'Unknown model type {model}')
def get_embeddings(model, model_type) -> list[torch.nn.Module]:
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL :
return [model.model.embed_tokens]
elif model_type == OPT_MODEL:
return [model.model.decoder.embed_tokens, model.model.decoder.embed_positions]
else:
raise ValueError(f'Unknown model type {model_type}')
def get_transformer_layers(model, model_type):
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL:
return [layer for layer in model.model.layers]
elif model_type == OPT_MODEL:
return [layer for layer in model.model.decoder.layers]
else:
raise ValueError(f'Unknown model type {model_type}')
def get_lm_head(model, model_type):
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL:
return model.lm_head
elif model_type == OPT_MODEL:
return model.lm_head
else:
raise ValueError(f'Unknown model type {model_type}')
def get_pre_head_layernorm(model, model_type):
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL :
pre_head_layernorm = model.model.norm
assert isinstance(pre_head_layernorm,
(transformers.models.llama.modeling_llama.LlamaRMSNorm,transformers.models.qwen2.modeling_qwen2.Qwen2RMSNorm))
elif model_type == OPT_MODEL:
pre_head_layernorm = model.model.decoder.final_layer_norm
assert pre_head_layernorm is not None
else:
raise ValueError(f'Unknown model type {model_type}')
return pre_head_layernorm
def fuse_ln_linear(layernorm: torch.nn.Module, linear_layers: typing.Iterable[torch.nn.Linear]) -> None:
"""
fuse the linear operations in Layernorm into the adjacent linear blocks.
"""
for linear in linear_layers:
linear_dtype = linear.weight.dtype
# Calculating new weight and bias
W_ = linear.weight.data.double()
linear.weight.data = (W_ * layernorm.weight.double()).to(linear_dtype)
if hasattr(layernorm, 'bias'):
if linear.bias is None:
linear.bias = torch.nn.Parameter(torch.zeros(linear.out_features, dtype=torch.float64))
linear.bias.data = linear.bias.data.double() + torch.matmul(W_, layernorm.bias.double())
linear.bias.data = linear.bias.data.to(linear_dtype)
def bake_mean_into_linear(linear: torch.nn.Linear) -> None:
"""
This function takes a linear layer and subtracts the means from the
weights and biases. This will result in the linear layer performing
the mean substitution which is usually done inside layernorm.
"""
linear_dtype = linear.weight.dtype
W_ = linear.weight.data.double()
linear.weight.data = W_ - W_.mean(dim=-2, keepdim=True)
linear.weight.data = linear.weight.data.to(linear_dtype)
if linear.bias is not None:
b_ = linear.bias.data.double()
linear.bias.data = b_ - b_.mean()
linear.bias.data = linear.bias.data.to(linear_dtype)
def fuse_layer_norms(model):
model_type = get_model_type(model)
kwargs = {'model': model, 'model_type': model_type}
# Embedding fusion
for W in get_embeddings(**kwargs):
W_ = W.weight.data.double()
W.weight.data = (W_ - W_.mean(dim=-1, keepdim=True)).to(W.weight.data.dtype)
layers = get_transformer_layers(**kwargs)
# Fuse the linear operations in Layernorm into the adjacent linear blocks.
for layer in layers:
# fuse the input layernorms into the linear layers
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL :
fuse_ln_linear(layer.post_attention_layernorm, [layer.mlp.up_proj, layer.mlp.gate_proj])
fuse_ln_linear(layer.input_layernorm, [layer.self_attn.q_proj, layer.self_attn.k_proj, layer.self_attn.v_proj])
elif model_type == OPT_MODEL:
fuse_ln_linear(layer.self_attn_layer_norm, [layer.self_attn.q_proj, layer.self_attn.k_proj, layer.self_attn.v_proj])
fuse_ln_linear(layer.final_layer_norm, [layer.fc1])
else:
raise ValueError(f'Unknown model type {model_type}')
if model_type == OPT_MODEL:
bake_mean_into_linear(layer.self_attn.out_proj)
bake_mean_into_linear(layer.fc2)
fuse_ln_linear(get_pre_head_layernorm(**kwargs), [get_lm_head(**kwargs)])
model_type2_NORM = {
LLAMA_MODEL: transformers.models.llama.modeling_llama.LlamaRMSNorm,
QWEN_MODEL: transformers.models.qwen2.modeling_qwen2.Qwen2RMSNorm,
OPT_MODEL: torch.nn.LayerNorm
}
replace_modules(
model,
model_type2_NORM[model_type],
lambda ori_norm: RMSN(model.config.hidden_size,ori_norm),
replace_layers=False,
)
def get_orthogonal_matrix(size, mode, device=DEV):
return random_hadamard_matrix(size, device)
def rotate_embeddings(model, Q: torch.Tensor) -> None:
# Rotate the embeddings.
model_type = model_type_extractor(model)
for W in get_embeddings(model, model_type):
dtype = W.weight.data.dtype
W_ = W.weight.data.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(W_, Q).to(device="cpu", dtype=dtype)
def rotate_attention_inputs(layer, Q, model_type) -> None:
# Rotate the WQ, WK and WV matrices of the self-attention layer.
for W in [layer.self_attn.q_proj, layer.self_attn.k_proj, layer.self_attn.v_proj]:
dtype = W.weight.dtype
W_ = W.weight.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(W_, Q).to(device="cpu", dtype=dtype)
def rotate_attention_output(layer, Q, model_type) -> None:
# Rotate output matrix of the self-attention layer.
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL:
W = layer.self_attn.o_proj
elif model_type == OPT_MODEL:
W = layer.self_attn.out_proj
else:
raise ValueError(f'Unknown model type {model_type}')
dtype = W.weight.data.dtype
W_ = W.weight.data.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(Q.T, W_).to(device="cpu", dtype=dtype)
if W.bias is not None:
b = W.bias.data.to(device=DEV, dtype=torch.float64)
W.bias.data = torch.matmul(Q.T, b).to(device="cpu", dtype=dtype)
def rotate_mlp_input(layer, Q, model_type):
# Rotate the MLP input weights.
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL:
mlp_inputs = [layer.mlp.up_proj, layer.mlp.gate_proj]
elif model_type == OPT_MODEL:
mlp_inputs = [layer.fc1]
else:
raise ValueError(f'Unknown model type {model_type}')
for W in mlp_inputs:
dtype = W.weight.dtype
W_ = W.weight.data.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(W_, Q).to(device="cpu", dtype=dtype)
def rotate_mlp_output(layer, Q, model_type,add_online_rotate=True):
# Rotate the MLP output weights and bias.
if model_type == LLAMA_MODEL or model_type == QWEN_MODEL:
W = layer.mlp.down_proj
elif model_type == OPT_MODEL:
W = layer.fc2
else:
raise ValueError(f'Unknown model type {model_type}')
dtype = W.weight.data.dtype
W_ = W.weight.data.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(Q.T, W_).to(device="cpu", dtype=dtype)
if W.bias is not None:
b = W.bias.data.to(device=DEV, dtype=torch.float64)
W.bias.data = torch.matmul(Q.T, b).to(device="cpu", dtype=dtype)
def rotate_head(model, Q: torch.Tensor) -> None:
# Rotate the head.
W = get_lm_head(model, model_type=model_type_extractor(model))
dtype = W.weight.data.dtype
W_ = W.weight.data.to(device=DEV, dtype=torch.float64)
W.weight.data = torch.matmul(W_, Q).to(device="cpu", dtype=dtype)
@torch.inference_mode()
def rotate_model(model, args,add_online_rotate=True):
Q = get_orthogonal_matrix(model.config.hidden_size,
args.rotate_mode)
config = model.config
num_heads = config.num_attention_heads
model_dim = config.hidden_size
head_dim = model_dim // num_heads
model_type = model_type_extractor(model)
rotate_embeddings(model, Q)
rotate_head(model, Q)
utils.cleanup_memory()
layers = get_transformer_layers(model, model_type=model_type)
for idx, layer in enumerate(tqdm.tqdm(layers, unit="layer", desc="Rotating")):
rotate_attention_inputs(layers[idx], Q, model_type)
rotate_attention_output(layers[idx], Q, model_type)
rotate_mlp_input(layers[idx], Q, model_type)
rotate_mlp_output(layers[idx], Q, model_type,add_online_rotate)
def get_rotate_model(model,save_path):
@dataclass
class parm():
rotate_mode:str
save_qmodel_path:str
fp32_had:bool
args = parm("hadamard",save_path,False)
model.to("cpu")
model.eval()
fuse_layer_norms(model)
rotate_model(model, args,add_online_rotate=False)
save_dict = {}
save_dict["model"] = model.state_dict()
torch.save(save_dict, args.save_qmodel_path)
return save_dict |