Bc-AI/Worker-Sam-z-api
0
1"""2SAM-Z-1 Distributed Worker Node v4.03Optimized for distributed gen/decode pipeline4"""5 6from fastapi import FastAPI, HTTPException7from fastapi.responses import StreamingResponse, HTMLResponse8from pydantic import BaseModel9import tensorflow as tf10import keras11from huggingface_hub import hf_hub_download12import json13import os14from tokenizers import Tokenizer15import numpy as np16import time17from typing import List, Optional18import asyncio19 20app = FastAPI(title="SAM-Z-1 Distributed Worker", version="4.0.0")21 22# ============================================================================23# Model Architecture24# ============================================================================25 26@keras.saving.register_keras_serializable()27class RotaryEmbedding(keras.layers.Layer):28 def __init__(self, dim, max_len=2048, theta=10000, **kwargs):29 super().__init__(**kwargs)30 self.dim = dim31 self.max_len = max_len32 self.theta = theta33 self.built_cache = False34 35 def build(self, input_shape):36 super().build(input_shape)37 38 def _build_cache(self):39 if not self.built_cache:40 inv_freq = 1.0 / (self.theta ** (tf.range(0, self.dim, 2, dtype=tf.float32) / self.dim))41 t = tf.range(self.max_len, dtype=tf.float32)42 freqs = tf.einsum("i,j->ij", t, inv_freq)43 emb = tf.concat([freqs, freqs], axis=-1)44 45 self.cos_cached = tf.constant(np.cos(emb.numpy()), dtype=tf.float32)46 self.sin_cached = tf.constant(np.sin(emb.numpy()), dtype=tf.float32)47 self.built_cache = True48 49 def rotate_half(self, x):50 x1, x2 = tf.split(x, 2, axis=-1)51 return tf.concat([-x2, x1], axis=-1)52 53 def call(self, q, k):54 self._build_cache()55 seq_len = tf.shape(q)[2]56 dtype = q.dtype57 cos = tf.cast(self.cos_cached[:seq_len, :], dtype)[None, None, :, :]58 sin = tf.cast(self.sin_cached[:seq_len, :], dtype)[None, None, :, :]59 60 q_rotated = (q * cos) + (self.rotate_half(q) * sin)61 k_rotated = (k * cos) + (self.rotate_half(k) * sin)62 63 return q_rotated, k_rotated64 65 def get_config(self):66 config = super().get_config()67 config.update({"dim": self.dim, "max_len": self.max_len, "theta": self.theta})68 return config69 70@keras.saving.register_keras_serializable()71class RMSNorm(keras.layers.Layer):72 def __init__(self, epsilon=1e-5, **kwargs):73 super().__init__(**kwargs)74 self.epsilon = epsilon75 76 def build(self, input_shape):77 self.scale = self.add_weight(name="scale", shape=(input_shape[-1],), initializer="ones")78 79 def call(self, x):80 variance = tf.reduce_mean(tf.square(x), axis=-1, keepdims=True)81 return x * tf.math.rsqrt(variance + self.epsilon) * self.scale82 83 def get_config(self):84 config = super().get_config()85 config.update({"epsilon": self.epsilon})86 return config87 88@keras.saving.register_keras_serializable()89class TransformerBlock(keras.layers.Layer):90 def __init__(self, d_model, n_heads, ff_dim, dropout, max_len, rope_theta, layer_idx=0, **kwargs):91 super().__init__(**kwargs)92 self.d_model = d_model93 self.n_heads = n_heads94 self.ff_dim = ff_dim95 self.dropout_rate = dropout96 self.max_len = max_len97 self.rope_theta = rope_theta98 self.head_dim = d_model // n_heads99 self.layer_idx = layer_idx100 101 self.pre_attn_norm = RMSNorm()102 self.pre_ffn_norm = RMSNorm()103 104 self.q_proj = keras.layers.Dense(d_model, use_bias=False, name="q_proj")105 self.k_proj = keras.layers.Dense(d_model, use_bias=False, name="k_proj")106 self.v_proj = keras.layers.Dense(d_model, use_bias=False, name="v_proj")107 self.out_proj = keras.layers.Dense(d_model, use_bias=False, name="o_proj")108 109 self.rope = RotaryEmbedding(self.head_dim, max_len=max_len, theta=rope_theta)110 111 self.gate_proj = keras.layers.Dense(ff_dim, use_bias=False, name="gate_proj")112 self.up_proj = keras.layers.Dense(ff_dim, use_bias=False, name="up_proj")113 self.down_proj = keras.layers.Dense(d_model, use_bias=False, name="down_proj")114 115 self.dropout = keras.layers.Dropout(dropout)116 117 def call(self, x, training=None):118 B, T, D = tf.shape(x)[0], tf.shape(x)[1], self.d_model119 dtype = x.dtype120 121 res = x122 y = self.pre_attn_norm(x)123 124 q = tf.transpose(tf.reshape(self.q_proj(y), [B, T, self.n_heads, self.head_dim]), [0, 2, 1, 3])125 k = tf.transpose(tf.reshape(self.k_proj(y), [B, T, self.n_heads, self.head_dim]), [0, 2, 1, 3])126 v = tf.transpose(tf.reshape(self.v_proj(y), [B, T, self.n_heads, self.head_dim]), [0, 2, 1, 3])127 128 q, k = self.rope(q, k)129 130 scores = tf.matmul(q, k, transpose_b=True) / tf.sqrt(tf.cast(self.head_dim, dtype))131 mask = tf.where(132 tf.linalg.band_part(tf.ones([T, T], dtype=dtype), -1, 0) == 0,133 tf.constant(-1e9, dtype=dtype),134 tf.constant(0.0, dtype=dtype)135 )136 scores += mask137 attn = tf.matmul(tf.nn.softmax(scores, axis=-1), v)138 139 attn = tf.reshape(tf.transpose(attn, [0, 2, 1, 3]), [B, T, D])140 x = res + self.dropout(self.out_proj(attn), training=training)141 142 res = x143 y = self.pre_ffn_norm(x)144 ffn = self.down_proj(keras.activations.silu(self.gate_proj(y)) * self.up_proj(y))145 146 return res + self.dropout(ffn, training=training)147 148 def get_config(self):149 config = super().get_config()150 config.update({151 "d_model": self.d_model,152 "n_heads": self.n_heads,153 "ff_dim": self.ff_dim,154 "dropout": self.dropout_rate,155 "max_len": self.max_len,156 "rope_theta": self.rope_theta,157 "layer_idx": self.layer_idx158 })159 return config160 161@keras.saving.register_keras_serializable()162class SAM1Model(keras.Model):163 def __init__(self, **kwargs):164 super().__init__()165 if 'config' in kwargs and isinstance(kwargs['config'], dict):166 self.cfg = kwargs['config']167 elif 'vocab_size' in kwargs:168 self.cfg = kwargs169 else:170 self.cfg = kwargs.get('cfg', kwargs)171 172 self.embed = keras.layers.Embedding(self.cfg['vocab_size'], self.cfg['d_model'], name="embed_tokens")173 174 ff_dim = int(self.cfg['d_model'] * self.cfg['ff_mult'])175 block_args = {176 'd_model': self.cfg['d_model'],177 'n_heads': self.cfg['n_heads'],178 'ff_dim': ff_dim,179 'dropout': self.cfg['dropout'],180 'max_len': self.cfg['max_len'],181 'rope_theta': self.cfg['rope_theta']182 }183 184 self.blocks = []185 for i in range(self.cfg['n_layers']):186 block = TransformerBlock(name=f"block_{i}", layer_idx=i, **block_args)187 self.blocks.append(block)188 189 self.norm = RMSNorm(name="final_norm")190 self.lm_head = keras.layers.Dense(self.cfg['vocab_size'], use_bias=False, name="lm_head")191 192 def call(self, input_ids, training=None):193 x = self.embed(input_ids)194 for block in self.blocks:195 x = block(x, training=training)196 return self.lm_head(self.norm(x))197 198 def get_config(self):199 base_config = super().get_config()200 base_config['config'] = self.cfg201 return base_config202 203# ============================================================================204# Global State205# ============================================================================206 207model = None208tokenizer = None209config = None210eos_token_id = None211fast_forward = None212 213MODEL_REPO = "Smilyai-labs/Sam-Z-1-tensorflow"214CACHE_DIR = "./model_cache"215 216# Stats217worker_stats = {218 "total_requests": 0,219 "total_tokens": 0,220 "decode_requests": 0,221 "uptime_start": time.time()222}223 224# ============================================================================225# Request Models226# ============================================================================227 228class GenerateRequest(BaseModel):229 prompt: str230 max_tokens: int = 512231 temperature: float = 0.8232 top_k: int = 40233 top_p: float = 0.9234 repetition_penalty: float = 1.1235 stream: bool = False236 return_token_ids: bool = False237 238class ChatMessage(BaseModel):239 role: str240 content: str241 242class ChatRequest(BaseModel):243 messages: List[ChatMessage]244 max_tokens: int = 512245 temperature: float = 0.8246 top_k: int = 40247 top_p: float = 0.9248 repetition_penalty: float = 1.1249 stream: bool = False250 return_token_ids: bool = False251 252class DecodeRequest(BaseModel):253 token_ids: List[int]254 255class BatchDecodeRequest(BaseModel):256 batches: List[List[int]]257 258# ============================================================================259# Generation Functions260# ============================================================================261 262def generate_tokens(263 prompt: str,264 max_tokens: int = 512,265 temperature: float = 0.8,266 top_k: int = 40,267 top_p: float = 0.9,268 repetition_penalty: float = 1.1,269 return_token_ids: bool = False270):271 """Core generation - yields (token_id, token_text or None)"""272 global model, tokenizer, config, eos_token_id, fast_forward273 274 input_ids = [i for i in tokenizer.encode(prompt).ids if i != eos_token_id]275 276 if len(input_ids) == 0:277 return278 279 if len(input_ids) > config['max_position_embeddings'] - max_tokens:280 input_ids = input_ids[-(config['max_position_embeddings'] - max_tokens):]281 282 input_tensor = tf.constant([input_ids], dtype=tf.int32)283 token_freq = {}284 285 for step in range(max_tokens):286 logits = fast_forward(input_tensor)287 next_token_logits = logits[0, -1, :].numpy()288 289 next_token_logits = next_token_logits / temperature290 291 if repetition_penalty != 1.0:292 for token_id, freq in token_freq.items():293 if token_id < len(next_token_logits):294 next_token_logits[token_id] /= (repetition_penalty ** freq)295 296 if top_k > 0:297 top_k_indices = np.argpartition(next_token_logits, -top_k)[-top_k:]298 top_k_logits = next_token_logits[top_k_indices]299 top_k_probs = tf.nn.softmax(top_k_logits).numpy()300 301 if top_p < 1.0:302 sorted_indices = np.argsort(top_k_probs)[::-1]303 cumsum = np.cumsum(top_k_probs[sorted_indices])304 cutoff_idx = np.searchsorted(cumsum, top_p)305 nucleus_indices = sorted_indices[:cutoff_idx + 1]306 307 nucleus_logits = top_k_logits[nucleus_indices]308 nucleus_probs = tf.nn.softmax(nucleus_logits).numpy()309 310 sampled_idx = np.random.choice(len(nucleus_probs), p=nucleus_probs)311 next_token_id = int(top_k_indices[nucleus_indices[sampled_idx]])312 else:313 sampled_idx = np.random.choice(len(top_k_probs), p=top_k_probs)314 next_token_id = int(top_k_indices[sampled_idx])315 else:316 probs = tf.nn.softmax(next_token_logits).numpy()317 next_token_id = np.random.choice(len(probs), p=probs)318 319 if next_token_id == eos_token_id:320 break321 322 token_freq[next_token_id] = token_freq.get(next_token_id, 0) + 1323 324 if return_token_ids:325 yield (next_token_id, None)326 else:327 token_text = tokenizer.decode([next_token_id])328 yield (next_token_id, token_text)329 330 input_tensor = tf.concat([input_tensor, [[next_token_id]]], axis=1)331 332 if input_tensor.shape[1] > config['max_position_embeddings']:333 input_tensor = input_tensor[:, -config['max_position_embeddings']:]334 335def format_chat_prompt(messages: List[ChatMessage]) -> str:336 prompt = ""337 for msg in messages:338 if msg.role == "user":339 prompt += f"<|im_start|>user\n{msg.content}<|im_end|>\n"340 elif msg.role == "assistant":341 prompt += f"<|im_start|>assistant\n{msg.content}<|im_end|>\n"342 343 prompt += "<|im_start|>assistant\n"344 return prompt345 346# ============================================================================347# Status Page348# ============================================================================349 350@app.get("/", response_class=HTMLResponse)351async def status_page():352 """Worker status page"""353 return """354<!DOCTYPE html>355<html>356<head>357 <title>SAM-Z-1 Worker Node</title>358 <style>359 * { margin: 0; padding: 0; box-sizing: border-box; }360 body {361 font-family: 'Courier New', monospace;362 background: linear-gradient(135deg, #1a1f3a 0%, #0a0e27 100%);363 color: #00bfff;364 padding: 20px;365 min-height: 100vh;366 }367 .container {368 max-width: 900px;369 margin: 0 auto;370 }371 .header {372 text-align: center;373 padding: 30px;374 background: rgba(0, 191, 255, 0.1);375 border: 2px solid #00bfff;376 border-radius: 10px;377 margin-bottom: 30px;378 box-shadow: 0 0 20px rgba(0, 191, 255, 0.3);379 }380 .header h1 {381 font-size: 2.5em;382 text-transform: uppercase;383 letter-spacing: 3px;384 animation: glow 2s ease-in-out infinite alternate;385 }386 @keyframes glow {387 from { text-shadow: 0 0 10px #00bfff; }388 to { text-shadow: 0 0 20px #00bfff, 0 0 30px #00bfff; }389 }390 .badge {391 display: inline-block;392 padding: 5px 15px;393 border-radius: 15px;394 font-size: 0.9em;395 margin-top: 10px;396 }397 .badge-ready {398 background: rgba(0, 255, 136, 0.2);399 border: 1px solid #00ff88;400 color: #00ff88;401 }402 .badge-loading {403 background: rgba(255, 165, 0, 0.2);404 border: 1px solid #ffa500;405 color: #ffa500;406 }407 .stats-grid {408 display: grid;409 grid-template-columns: repeat(auto-fit, minmax(200px, 1fr));410 gap: 20px;411 margin-bottom: 30px;412 }413 .stat-card {414 background: rgba(0, 191, 255, 0.05);415 border: 1px solid #00bfff;416 border-radius: 8px;417 padding: 20px;418 text-align: center;419 }420 .stat-label {421 font-size: 0.8em;422 opacity: 0.7;423 text-transform: uppercase;424 margin-bottom: 10px;425 }426 .stat-value {427 font-size: 2em;428 font-weight: bold;429 }430 .features {431 background: rgba(0, 191, 255, 0.05);432 border: 1px solid #00bfff;433 border-radius: 8px;434 padding: 20px;435 }436 .features h3 {437 margin-bottom: 15px;438 }439 .feature-list {440 list-style: none;441 padding: 0;442 }443 .feature-list li {444 padding: 10px;445 margin: 5px 0;446 background: rgba(0, 191, 255, 0.1);447 border-radius: 5px;448 }449 .feature-list li:before {450 content: "โก ";451 color: #00ff88;452 }453 .timestamp {454 text-align: center;455 margin-top: 20px;456 opacity: 0.5;457 }458 </style>459</head>460<body>461 <div class="container">462 <div class="header">463 <h1>โ๏ธ WORKER NODE โ๏ธ</h1>464 <div>SAM-Z-1 Distributed Worker v4.0</div>465 <div class="badge" id="status-badge">CHECKING STATUS...</div>466 </div>467 468 <div class="stats-grid" id="stats">469 <div class="stat-card">470 <div class="stat-label">Total Requests</div>471 <div class="stat-value" id="total-req">--</div>472 </div>473 <div class="stat-card">474 <div class="stat-label">Total Tokens</div>475 <div class="stat-value" id="total-tokens">--</div>476 </div>477 <div class="stat-card">478 <div class="stat-label">Decode Requests</div>479 <div class="stat-value" id="decode-req">--</div>480 </div>481 <div class="stat-card">482 <div class="stat-label">Uptime</div>483 <div class="stat-value" id="uptime">--</div>484 </div>485 </div>486 487 <div class="features">488 <h3>๐ CAPABILITIES</h3>489 <ul class="feature-list">490 <li>Full Text Generation</li>491 <li>Token-Only Mode (for distributed pipeline)</li>492 <li>High-Speed Batch Decoding</li>493 <li>Chat Completion</li>494 <li>Streaming & Non-Streaming</li>495 </ul>496 </div>497 498 <div class="timestamp" id="timestamp">Initializing...</div>499 </div>500 501 <script>502 async function updateStats() {503 try {504 const response = await fetch('/health');505 const data = await response.json();506 507 const badge = document.getElementById('status-badge');508 if (data.model_loaded) {509 badge.textContent = 'โ
READY FOR INFERENCE';510 badge.className = 'badge badge-ready';511 } else {512 badge.textContent = 'โณ LOADING MODEL...';513 badge.className = 'badge badge-loading';514 }515 516 // Fetch stats517 const statsRes = await fetch('/stats');518 const stats = await statsRes.json();519 520 document.getElementById('total-req').textContent = stats.total_requests;521 document.getElementById('total-tokens').textContent = stats.total_tokens;522 document.getElementById('decode-req').textContent = stats.decode_requests;523 524 const uptime = Math.floor(stats.uptime);525 const h = Math.floor(uptime / 3600);526 const m = Math.floor((uptime % 3600) / 60);527 const s = uptime % 60;528 document.getElementById('uptime').textContent = `${h}h ${m}m ${s}s`;529 530 document.getElementById('timestamp').textContent = 531 `Last update: ${new Date().toLocaleTimeString()}`;532 } catch (e) {533 console.error('Failed to update stats:', e);534 }535 }536 537 // Update every second538 setInterval(updateStats, 1000);539 updateStats();540 </script>541</body>542</html>543 """544 545# ============================================================================546# API Endpoints547# ============================================================================548 549@app.get("/health")550async def health():551 return {552 "status": "healthy" if model is not None else "loading",553 "model_loaded": model is not None554 }555 556@app.get("/stats")557async def stats():558 uptime = time.time() - worker_stats["uptime_start"]559 return {560 "total_requests": worker_stats["total_requests"],561 "total_tokens": worker_stats["total_tokens"],562 "decode_requests": worker_stats["decode_requests"],563 "uptime": uptime,564 "tokens_per_second": worker_stats["total_tokens"] / uptime if uptime > 0 else 0565 }566 567@app.post("/decode")568async def decode(request: DecodeRequest):569 """Fast single decode"""570 if tokenizer is None:571 raise HTTPException(status_code=503, detail="Tokenizer not loaded")572 573 try:574 worker_stats["decode_requests"] += 1575 text = tokenizer.decode(request.token_ids)576 return {"text": text}577 except Exception as e:578 raise HTTPException(status_code=500, detail=f"Decode error: {str(e)}")579 580@app.post("/decode/batch")581async def batch_decode(request: BatchDecodeRequest):582 """Optimized batch decoding for distributed pipeline"""583 if tokenizer is None:584 raise HTTPException(status_code=503, detail="Tokenizer not loaded")585 586 try:587 worker_stats["decode_requests"] += len(request.batches)588 results = [tokenizer.decode(batch) for batch in request.batches]589 return {"texts": results}590 except Exception as e:591 raise HTTPException(status_code=500, detail=f"Batch decode error: {str(e)}")592 593@app.post("/generate")594async def generate(request: GenerateRequest):595 """Generate text"""596 if model is None:597 raise HTTPException(status_code=503, detail="Model not loaded")598 599 worker_stats["total_requests"] += 1600 start_time = time.time()601 602 if request.stream:603 async def stream_tokens():604 generated_text = ""605 token_count = 0606 607 try:608 for token_id, token_text in generate_tokens(609 request.prompt,610 max_tokens=request.max_tokens,611 temperature=request.temperature,612 top_k=request.top_k,613 top_p=request.top_p,614 repetition_penalty=request.repetition_penalty,615 return_token_ids=request.return_token_ids616 ):617 token_count += 1618 worker_stats["total_tokens"] += 1619 620 if request.return_token_ids:621 yield f"data: {json.dumps({'token_id': token_id})}\n\n"622 else:623 generated_text += token_text624 yield f"data: {json.dumps({'text': token_text, 'total': generated_text})}\n\n"625 626 await asyncio.sleep(0.001)627 628 elapsed = time.time() - start_time629 yield f"data: {json.dumps({'done': True, 'tokens': token_count, 'time': elapsed})}\n\n"630 631 except Exception as e:632 yield f"data: {json.dumps({'error': str(e)})}\n\n"633 634 return StreamingResponse(stream_tokens(), media_type="text/event-stream")635 636 else:637 generated_text = ""638 token_count = 0639 640 try:641 for token_id, token_text in generate_tokens(642 request.prompt,643 max_tokens=request.max_tokens,644 temperature=request.temperature,645 top_k=request.top_k,646 top_p=request.top_p,647 repetition_penalty=request.repetition_penalty,648 return_token_ids=request.return_token_ids649 ):650 if not request.return_token_ids:651 generated_text += token_text652 token_count += 1653 worker_stats["total_tokens"] += 1654 655 elapsed = time.time() - start_time656 657 return {658 "text": generated_text,659 "tokens": token_count,660 "time": elapsed,661 "tokens_per_second": token_count / elapsed if elapsed > 0 else 0662 }663 664 except Exception as e:665 raise HTTPException(status_code=500, detail=f"Generation error: {str(e)}")666 667@app.post("/chat")668async def chat(request: ChatRequest):669 """Chat completion"""670 if model is None:671 raise HTTPException(status_code=503, detail="Model not loaded")672 673 worker_stats["total_requests"] += 1674 prompt = format_chat_prompt(request.messages)675 start_time = time.time()676 677 if request.stream:678 async def stream_tokens():679 generated_text = ""680 token_count = 0681 682 try:683 for token_id, token_text in generate_tokens(684 prompt,685 max_tokens=request.max_tokens,686 temperature=request.temperature,687 top_k=request.top_k,688 top_p=request.top_p,689 repetition_penalty=request.repetition_penalty,690 return_token_ids=request.return_token_ids691 ):692 token_count += 1693 worker_stats["total_tokens"] += 1694 695 if request.return_token_ids:696 yield f"data: {json.dumps({'token_id': token_id})}\n\n"697 else:698 generated_text += token_text699 700 if "<|im_end|>" in generated_text:701 generated_text = generated_text.split("<|im_end|>")[0]702 break703 704 yield f"data: {json.dumps({'delta': token_text, 'content': generated_text})}\n\n"705 706 await asyncio.sleep(0.001)707 708 elapsed = time.time() - start_time709 yield f"data: {json.dumps({'done': True, 'tokens': token_count, 'time': elapsed})}\n\n"710 711 except Exception as e:712 yield f"data: {json.dumps({'error': str(e)})}\n\n"713 714 return StreamingResponse(stream_tokens(), media_type="text/event-stream")715 716 else:717 generated_text = ""718 token_count = 0719 720 try:721 for token_id, token_text in generate_tokens(722 prompt,723 max_tokens=request.max_tokens,724 temperature=request.temperature,725 top_k=request.top_k,726 top_p=request.top_p,727 repetition_penalty=request.repetition_penalty,728 return_token_ids=request.return_token_ids729 ):730 if not request.return_token_ids:731 generated_text += token_text732 733 if "<|im_end|>" in generated_text:734 generated_text = generated_text.split("<|im_end|>")[0]735 break736 737 token_count += 1738 worker_stats["total_tokens"] += 1739 740 elapsed = time.time() - start_time741 742 return {743 "message": {744 "role": "assistant",745 "content": generated_text.strip()746 },747 "tokens": token_count,748 "time": elapsed,749 "tokens_per_second": token_count / elapsed if elapsed > 0 else 0750 }751 752 except Exception as e:753 raise HTTPException(status_code=500, detail=f"Generation error: {str(e)}")754 755# ============================================================================756# Model Loading757# ============================================================================758 759@app.on_event("startup")760async def load_model():761 global model, tokenizer, config, eos_token_id, fast_forward762 763 print("๐ Loading SAM-Z-1 Model...")764 765 try:766 config_path = hf_hub_download(MODEL_REPO, "config.json", cache_dir=CACHE_DIR)767 768 try:769 weights_path = hf_hub_download(MODEL_REPO, "ckpt.weights.h5", cache_dir=CACHE_DIR)770 print("โ
Found checkpoint weights")771 use_checkpoint = True772 except:773 print("โ ๏ธ Checkpoint not found, using model.keras")774 model_path = hf_hub_download(MODEL_REPO, "model.keras", cache_dir=CACHE_DIR)775 use_checkpoint = False776 777 with open(config_path, 'r') as f:778 config = json.load(f)779 780 print(f"๐ฆ Config loaded: {config['num_hidden_layers']} layers")781 782 print("๐ฆ Creating tokenizer...")783 from transformers import AutoTokenizer784 785 hf_tokenizer = AutoTokenizer.from_pretrained("gpt2")786 custom_tokens = ["<|im_start|>", "<|im_end|>", "<think>", "<think/>"]787 hf_tokenizer.add_special_tokens({"additional_special_tokens": custom_tokens})788 789 os.makedirs("./temp_tokenizer", exist_ok=True)790 hf_tokenizer.save_pretrained("./temp_tokenizer")791 tokenizer = Tokenizer.from_file("./temp_tokenizer/tokenizer.json")792 793 eos_token_id = config.get('eos_token_id', 50256)794 795 print(f"โ
Tokenizer ready: vocab size {tokenizer.get_vocab_size()}")796 797 print("๐ Loading model...")798 799 if use_checkpoint:800 model_config = {801 'vocab_size': config['vocab_size'],802 'd_model': config['hidden_size'],803 'n_layers': config['num_hidden_layers'],804 'n_heads': config['num_attention_heads'],805 'ff_mult': config['intermediate_size'] / config['hidden_size'],806 'max_len': config['max_position_embeddings'],807 'dropout': 0.1,808 'rope_theta': config['rope_theta']809 }810 811 model = SAM1Model(config=model_config)812 dummy_input = tf.zeros((1, config['max_position_embeddings']), dtype=tf.int32)813 _ = model(dummy_input, training=False)814 815 print(f"โ
Architecture built: {model.count_params():,} parameters")816 817 model.load_weights(weights_path)818 print("โ
Weights loaded!")819 820 else:821 model = keras.models.load_model(model_path, compile=False)822 print("โ
Model loaded!")823 824 @tf.function(reduce_retracing=True)825 def optimized_forward(input_tensor):826 return model(input_tensor, training=False)827 828 fast_forward = optimized_forward829 830 print("โ
SAM-Z-1 Distributed Worker ready! ๐")831 print("๐ฅ Features enabled:")832 print(" - Full text generation")833 print(" - Token-only mode (distributed pipeline)")834 print(" - Batch decoding optimization")835 print(" - Streaming support")836 837 except Exception as e:838 print(f"โ Failed to load model: {e}")839 import traceback840 traceback.print_exc()841 raise842 843# ============================================================================844# Launch845# ============================================================================846 847if __name__ == "__main__":848 import uvicorn849 uvicorn.run(850 app,851 host="0.0.0.0",852 port=7860,853 log_level="info"854 )