import gradio as gr import torch import pickle import torch.nn as nn import torch.nn.functional as F import math import urllib.request from bs4 import BeautifulSoup from googlesearch import search # นำเข้าเครื่องมือเปิดเว็บเซอร์วิสสากลและเครือข่าย from fastapi import FastAPI, Request from fastapi.middleware.cors import CORSMiddleware import uvicorn # ========================================== # 0. คลาสตัวตัดคำระดับอักขระดั้งเดิม # ========================================== class CharTokenizer: def __init__(self, text): self.chars = sorted(list(set(text))) self.vocab_size = len(self.chars) self.stoi = { ch:i for i,ch in enumerate(self.chars) } self.itos = { i:ch for i,ch in enumerate(self.chars) } def encode(self, s): return [self.stoi[c] for c in s if c in self.stoi] def decode(self, l): return ''.join([self.itos[i] for i in l if i in self.itos]) import __main__ __main__.CharTokenizer = CharTokenizer # ========================================== # 1. โครงสร้างสถาปัตยกรรมโมเดลขั้นสูง (เหมือนเดิม 100%) # ========================================== n_embd = 768 block_size = 256 n_heads = 12 n_kv_heads = 4 n_layers = 8 ffn_hidden_dim = 2048 class GemmaRMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.zeros(dim)) def forward(self, x): variance = x.pow(2).mean(-1, keepdim=True) return x * torch.rsqrt(variance + self.eps) * (1.0 + self.weight) class GemmaSwiGLU(nn.Module): def __init__(self, d_in: int, d_hidden: int): super().__init__() self.gate_proj = nn.Linear(d_in, d_hidden, bias=False) self.up_proj = nn.Linear(d_in, d_hidden, bias=False) self.down_proj = nn.Linear(d_hidden, d_in, bias=False) def forward(self, x): return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class GemmaRotaryEmbedding(nn.Module): def __init__(self, dim, max_seq_len=2048, theta=10000.0): super().__init__() self.dim = dim inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) t = torch.arange(max_seq_len, dtype=torch.float32) freqs = torch.outer(t, self.inv_freq) emb = torch.cat((freqs, freqs), dim=-1) self.register_buffer("cos_cached", emb.cos(), persistent=False) self.register_buffer("sin_cached", emb.sin(), persistent=False) def forward(self, x, seq_len): return self.cos_cached[:seq_len, :], self.sin_cached[:seq_len, :] def rotate_half(x): x1 = x[..., :x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2:] return torch.cat((-x2, x1), dim=-1) def apply_rope(q, k, cos, sin): cos = cos.unsqueeze(0).unsqueeze(2) sin = sin.unsqueeze(0).unsqueeze(2) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed class GemmaGroupedAttention(nn.Module): def __init__(self): super().__init__() self.head_dim = n_embd // n_heads self.num_local_heads = n_heads self.num_local_kv_heads = n_kv_heads self.num_queries_per_kv = n_heads // n_kv_heads self.q_proj = nn.Linear(n_embd, n_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(n_embd, n_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(n_embd, n_kv_heads * self.head_dim, bias=False) self.o_proj = nn.Linear(n_heads * self.head_dim, n_embd, bias=False) self.rope = GemmaRotaryEmbedding(self.head_dim) self.register_buffer('tril', torch.tril(torch.ones(block_size, block_size))) def forward(self, x): B, T, C = x.shape q = self.q_proj(x).view(B, T, self.num_local_heads, self.head_dim) k = self.k_proj(x).view(B, T, self.num_local_kv_heads, self.head_dim) v = self.v_proj(x).view(B, T, self.num_local_kv_heads, self.head_dim) cos, sin = self.rope(q, T) q, k = apply_rope(q, k, cos, sin) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) if self.num_queries_per_kv > 1: k = k.repeat_interleave(self.num_queries_per_kv, dim=1) v = v.repeat_interleave(self.num_queries_per_kv, dim=1) scores = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim) scores = torch.tanh(scores / 10.0) * 10.0 scores = scores.masked_fill(self.tril[:T, :T] == 0, float('-inf')) attention_probs = F.softmax(scores, dim=-1) output = attention_probs @ v output = output.transpose(1, 2).contiguous().view(B, T, C) return self.o_proj(output) class GemmaDecoderBlock(nn.Module): def __init__(self): super().__init__() self.attn = GemmaGroupedAttention() self.ffn = GemmaSwiGLU(n_embd, ffn_hidden_dim) self.input_layernorm = GemmaRMSNorm(n_embd) self.post_attention_layernorm = GemmaRMSNorm(n_embd) def forward(self, x): x = x + self.attn(self.input_layernorm(x)) x = x + self.ffn(self.post_attention_layernorm(x)) return x class DeepTanzGemmaModel(nn.Module): def __init__(self, vocab_size): super().__init__() self.embed = nn.Embedding(vocab_size, n_embd) self.layers = nn.ModuleList([GemmaDecoderBlock() for _ in range(n_layers)]) self.norm = GemmaRMSNorm(n_embd) self.lm_head = nn.Linear(n_embd, vocab_size, bias=False) self.embed.weight = self.lm_head.weight def forward(self, idx): x = self.embed(idx) * math.sqrt(n_embd) for layer in self.layers: x = layer(x) x = self.norm(x) return self.lm_head(x) # โหลดระบบตัดคำศัพท์และไฟล์โมเดลที่ผ่านการเทรน with open('tokenizer.pkl', 'rb') as f: tokenizer = pickle.load(f) model = DeepTanzGemmaModel(tokenizer.vocab_size) state_dict = torch.load('advanced_gemini.pt', map_location=torch.device('cpu')) model.load_state_dict(state_dict) model.eval() # ========================================== # 2. ระบบดึงข้อมูลสด Google Search # ========================================== def fetch_google_knowledge(query): try: search_results = list(search(query, num_results=1, lang="th")) if not search_results: return "" target_url = search_results[0] req = urllib.request.Request(target_url, headers={'User-Agent': 'Mozilla/5.0'}) html = urllib.request.urlopen(req, timeout=5).read() soup = BeautifulSoup(html, 'html.parser') for script in soup(["script", "style"]): script.extract() text = soup.get_text() lines = (line.strip() for line in text.splitlines()) chunks = (phrase.strip() for line in lines for phrase in line.split(" ")) clean_text = " ".join(chunk for chunk in chunks if chunk) return clean_text[:180] except Exception: return "" # ========================================== # 3. ฟังก์ชันหลักสำหรับประมวลผลผ่าน API # ========================================== def api_predict(user_input): full_prompt = f"Q: {user_input}\nA: " idx = torch.tensor([tokenizer.encode(full_prompt)], dtype=torch.long) max_new_tokens = 180 generated_tokens = [] with torch.no_grad(): initial_logits = model(idx[:, -block_size:])[:, -1, :] probs = F.softmax(initial_logits, dim=-1) max_prob, _ = torch.max(probs, dim=-1) if max_prob.item() < 0.15: live_info = fetch_google_knowledge(user_input) if live_info: full_prompt = f"ข้อมูลเพิ่มเติมจากกูเกิ้ล: {live_info}\nQ: {user_input}\nA: " idx = torch.tensor([tokenizer.encode(full_prompt)], dtype=torch.long) else: return "ผมไม่สามารถตอบคำถามนี้ได้เนื่องจากไม่พบคลังข้อมูลบนระบบอินเทอร์เน็ตครับ" for _ in range(max_new_tokens): idx_cond = idx[:, -block_size:] with torch.no_grad(): logits = model(idx_cond)[:, -1, :] v, ix = torch.topk(logits, k=3) filtered_logits = torch.full_like(logits, -float('Inf')) filtered_logits.scatter_(1, ix, v) idx_next = torch.multinomial(F.softmax(filtered_logits, dim=-1), num_samples=1) idx = torch.cat((idx, idx_next), dim=1) generated_tokens.append(idx_next.item()) generated_text = tokenizer.decode(generated_tokens) if "\n" in generated_text or "Q:" in generated_text or "A:" in generated_text: clean_output = generated_text.replace("\n", "").replace("Q:", "").replace("A:", "").strip() return clean_output if clean_output else "ผมไม่สามารถหาข้อสรุปจากเนื้อหาได้ครับ" return tokenizer.decode(generated_tokens).strip() # ========================================== # 4. ประกอบร่างเซิร์ฟเวอร์ด้วย FastAPI และปลดบล็อก CORS # ========================================== app = FastAPI() # 🔓 เปิดประตูรับข้อมูลข้ามเว็บจากทุกที่แบบอิสระ หมดปัญหาบราวเซอร์บล็อก CORS app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # ประตูรับส่งสัญญาณข้ามเว็บที่พิกัด /api/chat @app.post("/api/chat") async def chat_endpoint(request: Request): payload = await request.json() user_input = payload.get("text", "") # ดึงคำตอบจากโครงข่ายหลัก ai_response = api_predict(user_input) return {"response": ai_response} # ผูกอินเตอร์เฟส Gradio เปล่าหน้าบ้านเพื่อรักษาระบบล็อกของ Spaces with gr.Blocks() as demo: gr.Markdown("### ⚡ Tanz Gemma API Backend via FastAPI is fully operational.") # เชื่อมโครงข่ายมิติคู่ขนาน app = gr.mount_gradio_app(app, demo, path="/") # 🔥 บล็อกแก้ทาง Exit code 0 และ Error 98: สั่งเปิดเซิร์ฟเวอร์ค้างไว้ตลอดกาลผ่านโมดูลหลักสากล if __name__ == "__main__": uvicorn.run("app:app", host="0.0.0.0", port=7860, reload=False, workers=1)