Views
No views yet
-1, 0, y 1 para mejorar la eficiencia en el cómputo sin perder precisión.
1from transformers import AutoModelForCausalLM, AutoTokenizer
2from transformers.models.llama.modeling_llama import *
3import torch
4from torch import nn
5import torch.nn.functional as F
6import coloredlogs
7import logging
8
9from utils.utils import count_parameters
10
11coloredlogs.install(level='INFO', fmt='%(asctime)s - %(levelname)s - %(message)s', logger=logging.getLogger())
12logger = logging.getLogger(__name__)
13
14
15
16
17HF_TOKEN = "tuclaveaqui"
18#model = "ejbejaranos/Bitnet-Llama3-from8BM-now2B"
19model = "ejbejaranos/Bitnet-Nous-Llama3-225M" ## Working
20
21# Load a pretrained BitNet model
22tokenizer = AutoTokenizer.from_pretrained(model)
23
24model = AutoModelForCausalLM.from_pretrained(
25 model,
26 token=HF_TOKEN
27)
28
29
30def count_parameters(model):
31 # Calculate the number of parameters in billions
32 num_params = sum(p.numel() for p in model.parameters() if p.requires_grad) / 10**9
33 print(f"Model size: {num_params:.3f}B parameters")
34 return int(num_params)
35
36
37
38# Establece el pad_token_id
39model.config.pad_token_id = tokenizer.eos_token_id
40
41def activation_quant(x):
42 scale = 127.0 / x.abs().max(dim=-1, keepdim=True).values.clamp_(min=1e-5)
43 y = (x * scale).round().clamp_(-128, 127)
44 y = y / scale
45 return y
46
47def weight_quant(w):
48 scale = 1.0 / w.abs().mean().clamp_(min=1e-5)
49 u = (w * scale).round().clamp_(-1, 1)
50 u = u / scale
51 return u
52
53class BitLinear(nn.Linear):
54 def forward(self, x):
55 w = self.weight # a weight tensor with shape [d, k]
56 x = x.to(w.device)
57 RMSNorm = LlamaRMSNorm(x.shape[-1]).to(w.device)
58 x_norm = RMSNorm(x)
59 x_quant = x_norm + (activation_quant(x_norm) - x_norm).detach()
60 w_quant = w + (weight_quant(w) - w).detach()
61 y = F.linear(x_quant, w_quant)
62 return y
63
64def convert_to_bitnet(model, copy_weights):
65 for name, module in model.named_modules():
66 if isinstance(module, LlamaSdpaAttention) or isinstance(module, LlamaMLP):
67 for child_name, child_module in module.named_children():
68 if isinstance(child_module, nn.Linear):
69 bitlinear = BitLinear(child_module.in_features, child_module.out_features, child_module.bias is not None).to(device="cuda:0")
70 if copy_weights:
71 bitlinear.weight = child_module.weight
72 if child_module.bias is not None:
73 bitlinear.bias = child_module.bias
74 setattr(module, child_name, bitlinear)
75 elif isinstance(module, LlamaDecoderLayer):
76 for child_name, child_module in module.named_children():
77 if isinstance(child_module, LlamaRMSNorm) and child_name == "input_layernorm":
78 setattr(module, child_name, nn.Identity().to(device="cuda:0"))
79
80convert_to_bitnet(model, copy_weights=True)
81model.to(device="cuda:0")
82
83
84logger.info(f"🔢 Number of parameters in the model after extracting weights: {count_parameters(model)}")
85logger.info(f"📏 Reduced model structure:\n{model}")
86
87
88
89
90
91prompt = "What is Machine Learning?"
92inputs = tokenizer(prompt, return_tensors="pt", padding=True, truncation=True).to(model.device)
93inputs['attention_mask'] = inputs['input_ids'] != model.config.pad_token_id
94
95generate_ids = model.generate(inputs.input_ids, attention_mask=inputs['attention_mask'], max_length=250)
96decoded_output = tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)
97
98print(decoded_output[0]) # Print the generated response
99