CoolFace
Apppublic

Bc-AI/Worker-Sam-z-api

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
app.py854 linesDownload Raw Back to root
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    )