Views
No views yet

pip install transformers torch1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4from huggingface_hub import hf_hub_download
5from safetensors.torch import load_file
6from transformers import PreTrainedTokenizerFast
7import os
8
9# ----------------------------
10# Config modello
11# ----------------------------
12VOCAB_SIZE = 1920
13D_MODEL = 240
14N_LAYERS = 6
15N_HEADS = 6
16D_HEAD = 40
17D_FF = 960
18N_CTX = 64
19
20device = "cuda" if torch.cuda.is_available() else "cpu"
21print(f"Usando device: {device}")
22
23# ----------------------------
24# Scarica tokenizer e modello da HF
25# ----------------------------
26repo_id = "Mattimax/PicoDAC"
27
28tokenizer_path = hf_hub_download(repo_id, "tokenizer.json")
29model_path = hf_hub_download(repo_id, "model.safetensors")
30# opzionale: scales_path = hf_hub_download(repo_id, "best/scales.json")
31
32# ----------------------------
33# Carica tokenizer
34# ----------------------------
35tokenizer = PreTrainedTokenizerFast(tokenizer_file=tokenizer_path)
36tokenizer.add_special_tokens({
37 "pad_token": "<PAD>",
38 "bos_token": "<BOS>",
39 "eos_token": "<EOS>",
40 "sep_token": "<SEP>",
41 "unk_token": "<UNK>"
42})
43
44# ----------------------------
45# RMSNorm
46# ----------------------------
47class RMSNorm(nn.Module):
48 def __init__(self, dim, eps=1e-8):
49 super().__init__()
50 self.eps = eps
51 self.weight = nn.Parameter(torch.ones(dim))
52
53 def forward(self, x):
54 norm = x.pow(2).mean(-1, keepdim=True).add(self.eps).sqrt()
55 return x / norm * self.weight
56
57# ----------------------------
58# Transformer Block
59# ----------------------------
60class TransformerBlock(nn.Module):
61 def __init__(self):
62 super().__init__()
63 self.ln1 = RMSNorm(D_MODEL)
64 self.ln2 = RMSNorm(D_MODEL)
65 self.attn = nn.MultiheadAttention(D_MODEL, N_HEADS, batch_first=True)
66 self.ffn = nn.Sequential(
67 nn.Linear(D_MODEL, D_FF),
68 nn.SiLU(),
69 nn.Linear(D_FF, D_MODEL)
70 )
71
72 def forward(self, x, mask=None):
73 attn_out, _ = self.attn(self.ln1(x), self.ln1(x), self.ln1(x), attn_mask=mask)
74 x = x + attn_out
75 x = x + self.ffn(self.ln2(x))
76 return x
77
78# ----------------------------
79# TinyGPT
80# ----------------------------
81class TinyGPT(nn.Module):
82 def __init__(self):
83 super().__init__()
84 self.tok_emb = nn.Embedding(VOCAB_SIZE, D_MODEL)
85 self.pos_emb = nn.Embedding(N_CTX, D_MODEL)
86 self.layers = nn.ModuleList([TransformerBlock() for _ in range(N_LAYERS)])
87 self.ln_f = RMSNorm(D_MODEL)
88 self.head = nn.Linear(D_MODEL, VOCAB_SIZE, bias=False)
89 self.head.weight = self.tok_emb.weight
90
91 def forward(self, idx):
92 B, T = idx.shape
93 pos = torch.arange(T, device=idx.device).unsqueeze(0)
94 x = self.tok_emb(idx) + self.pos_emb(pos)
95 mask = torch.triu(torch.ones(T, T, device=idx.device) * float('-inf'), diagonal=1)
96 for layer in self.layers:
97 x = layer(x, mask=mask)
98 x = self.ln_f(x)
99 logits = self.head(x)
100 return logits
101
102# ----------------------------
103# Carica pesi
104# ----------------------------
105state_dict = load_file(model_path)
106model = TinyGPT()
107model.load_state_dict(state_dict, strict=False)
108model.to(device)
109model.eval()
110
111# ----------------------------
112# Funzione generazione
113# ----------------------------
114def generate(prompt, max_new_tokens=50, temperature=0.7):
115 model_input = tokenizer.encode(prompt, return_tensors="pt").to(device)
116 for _ in range(max_new_tokens):
117 logits = model(model_input)
118 next_token_logits = logits[0, -1, :] / temperature
119 probs = F.softmax(next_token_logits, dim=-1)
120 next_token = torch.multinomial(probs, num_samples=1)
121 model_input = torch.cat([model_input, next_token.unsqueeze(0)], dim=1)
122 if next_token.item() == tokenizer.eos_token_id:
123 break
124 return tokenizer.decode(model_input[0], skip_special_tokens=True)
125
126# ----------------------------
127# Loop chat
128# ----------------------------
129print("Chat PicoDAC (scrivi 'exit' per uscire)")
130while True:
131 user_input = input("Tu: ").strip()
132 if user_input.lower() in ["exit", "quit"]:
133 print("Chiusura chat. Ciao!")
134 break
135 response = generate(user_input)
136 print(f"PicoDAC: {response}")max_length basso per mantenere la coerenza delle risposte.