Views
No views yet
CausalSelfAttentionFixed dinamis.vocab_size = 3000).torch.compile untuk penyusunan kernel C++ GPU yang super cepat saat training.model.pt: File matriks bobot (state_dict) PyTorch mentah berukuran ~349 MB.tokenizer.json: File konfigurasi BPE Tokenizer skala industri berbasis Rust.1import os
2import re
3import torch
4import torch.nn as nn
5from torch.nn import functional as F
6from tokenizers import Tokenizer
7from huggingface_hub import hf_hub_download
8
9# 1. Definisikan Kelas Arsitektur Model Kustom
10class CausalSelfAttentionFixed(nn.Module):
11 def __init__(self, d_model, n_heads):
12 super().__init__()
13 self.n_heads = n_heads
14 self.d_model = d_model
15 self.c_attn = nn.Linear(d_model, 3 * d_model)
16 self.c_proj = nn.Linear(d_model, d_model)
17
18 def forward(self, x):
19 B, T, C = x.size()
20 q, k, v = self.c_attn(x).split(self.d_model, dim=2)
21 k = k.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2)
22 q = q.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2)
23 v = v.view(B, T, self.n_heads, C // self.n_heads).transpose(1, 2)
24 att = (q @ k.transpose(-2, -1)) * (1.0 / (k.size(-1) ** 0.5))
25 mask = torch.tril(torch.ones(T, T, device=x.device)).view(1, 1, T, T)
26 att = att.masked_fill(mask == 0, float('-inf'))
27 att = F.softmax(att, dim=-1)
28 y = att @ v
29 return self.c_proj(y.transpose(1, 2).contiguous().view(B, T, C))
30
31class BlockFixed(nn.Module):
32 def __init__(self, d_model, n_heads):
33 super().__init__()
34 self.ln_1 = nn.LayerNorm(d_model)
35 self.attn = CausalSelfAttentionFixed(d_model, n_heads)
36 self.ln_2 = nn.LayerNorm(d_model)
37 self.mlp = nn.Sequential(nn.Linear(d_model, 4 * d_model), nn.GELU(), nn.Linear(4 * d_model, d_model))
38 def forward(self, x):
39 x = x + self.attn(self.ln_1(x))
40 return x + self.mlp(self.ln_2(x))
41
42class WarsaModelFixed(nn.Module):
43 def __init__(self, vocab_size, d_model=768, n_heads=12, n_layers=12, max_len=128):
44 super().__init__()
45 self.transformer = nn.ModuleDict(dict(
46 wte = nn.Embedding(vocab_size, d_model),
47 wpe = nn.Embedding(max_len, d_model),
48 h = nn.ModuleList([BlockFixed(d_model, n_heads) for _ in range(n_layers)]),
49 ln_f = nn.LayerNorm(d_model),
50 ))
51 self.lm_head = nn.Linear(d_model, vocab_size, bias=False)
52 self.transformer.wte.weight = self.lm_head.weight
53 def forward(self, idx, targets=None):
54 device = idx.device
55 b, t = idx.size()
56 pos = torch.arange(0, t, dtype=torch.long, device=device).unsqueeze(0)
57 x = self.transformer.wte(idx) + self.transformer.wpe(pos)
58 for block in self.transformer.h: x = block(x)
59 return self.lm_head(self.transformer.ln_f(x)), None
60
61# 2. Pipeline Inferensi Bersih
62def generate_warsa_clean(model, tokenizer, prompt, max_new_tokens=50, temperature=0.01):
63 model.eval()
64 device = next(model.parameters()).device
65 tokens = tokenizer.encode(prompt).ids
66 x = torch.tensor(tokens, dtype=torch.long, device=device).unsqueeze(0)
67 eos_id = tokenizer.token_to_id("<|endoftext|>")
68
69 for _ in range(max_new_tokens):
70 x_cond = x if x.size(1) <= 128 else x[:, -128:]
71 with torch.no_grad():
72 with torch.amp.autocast(device_type='cuda', dtype=torch.float16 if torch.cuda.is_available() else torch.float32):
73 logits, _ = model(x_cond)
74 logits = logits[:, -1, :] / temperature
75 next_token = torch.argmax(logits, dim=-1, keepdim=True)
76 x = torch.cat((x, next_token), dim=1)
77 if next_token.item() == eos_id: break
78
79 raw_output = tokenizer.decode(x[0].tolist())
80 response = raw_output.split("<|assistant|>")[-1].strip() if "<|assistant|>" in raw_output else raw_output.strip()
81 return re.sub(r'\s+([,.!?;:])', r'\1', response)
82
83# 3. Download & Jalankan Model
84REPO_ID = "dhiiitraaa/warsa-110m-industrial"
85device = "cuda" if torch.cuda.is_available() else "cpu"
86
87tokenizer_path = hf_hub_download(repo_id=REPO_ID, filename="tokenizer.json")
88model_path = hf_hub_download(repo_id=REPO_ID, filename="model.pt")
89
90tokenizer = Tokenizer.from_file(tokenizer_path)
91model = WarsaModelFixed(vocab_size=tokenizer.get_vocab_size())
92model.load_state_dict(torch.load(model_path, map_location=device))
93model.to(device)
94
95# Uji Coba Chat
96prompt = "<|user|>\nHai!<|assistant|>\n"
97print("Warsa:", generate_warsa_clean(model, tokenizer, prompt))