CoolFace
Apppublic

LeoHai123/SmaLLMPro

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
server.js134 linesDownload Raw Back to root
1const express = require('express');2const ort = require('onnxruntime-node');3const tiktoken = require('js-tiktoken');4const cors = require('cors');5const path = require('path'); // WICHTIG: Path Modul hinzufügen6 7const app = express();8app.use(cors());9app.use(express.json());10 11const enc = tiktoken.getEncoding("gpt2"); 12let session = null;13 14async function initModel() {15    console.log("--- DEBUG: DATEI-CHECK ---");16    const fs = require('fs');17    try {18        // Hier den Namen EXAKT anpassen:19        const modelPath = path.join(__dirname, 'model_124M_instruct_web_quant_base.onnx');20        21        console.log("Gesuchter Pfad:", modelPath);22        if (fs.existsSync(modelPath)) {23            session = await ort.InferenceSession.create(modelPath);24            console.log("Modell erfolgreich geladen!");25        } else {26            console.error("DATEI IMMER NOCH NICHT GEFUNDEN!");27        }28    } catch (e) {29        console.error("Fehler:", e.message);30    }31}32initModel();33 34app.post('/chat', async (req, res) => {35    if (!session) return res.status(503).json({ error: "Modell lädt noch..." });36 37    // WICHTIG: Variable MUSS hier drinnen neu erstellt werden für JEDEN Request38    let clientConnected = true; 39 40    // Wir hören auf das Response-Objekt41    res.on('close', () => {42        clientConnected = false;43        console.log("Verbindung geschlossen.");44    });45  46    const { prompt, temp, topK, maxLen, penalty } = req.body;47    res.setHeader('Content-Type', 'text/event-stream');48    res.setHeader('Cache-Control', 'no-cache');49 50    const formattedPrompt = `Instruction:\n${prompt}\n\nResponse:\n`;51    let tokens = enc.encode(formattedPrompt);52 53    // WICHTIG: Deine spezifische Vokabular-Größe54    const VOCAB_SIZE = 50304; 55 56    try {57        for (let i = 0; i < maxLen; i++) {58            // CHECK: Nur weitermachen, wenn der Client noch da ist59            if (!clientConnected) {60                console.log("Inferenz abgebrochen, da Client weg.");61                break; 62            }63          64            const ctx = tokens.slice(-1024);65            const inputData = BigInt64Array.from(ctx.map(x => BigInt(x)));66            const tensor = new ort.Tensor('int64', inputData, [1, ctx.length]);67 68            const results = await session.run({ input: tensor });69            const outputName = session.outputNames[0];70            71            // Wir nehmen exakt die letzten VOCAB_SIZE Werte72            const logits = Array.from(results[outputName].data.slice(-VOCAB_SIZE));73 74            // 1. Repetition Penalty (Exakt wie in deinem Original)75            if (penalty !== 1.0) {76                for (const token of tokens) {77                    if (token < VOCAB_SIZE) {78                        if (logits[token] > 0) {79                            logits[token] /= penalty;80                        } else {81                            logits[token] *= penalty;82                        }83                    }84                }85            }86 87            // 2. Sampling (Exakt wie in deinem Original)88            let scaledLogits = logits.map(l => l / temp);89            const maxLogit = Math.max(...scaledLogits);90            const exps = scaledLogits.map(l => Math.exp(l - maxLogit));91            const sumExps = exps.reduce((a, b) => a + b, 0);92            let probs = exps.map(e => e / sumExps);93 94            let indexedProbs = probs.map((p, i) => ({ p, i }));95            indexedProbs.sort((a, b) => b.p - a.p);96            indexedProbs = indexedProbs.slice(0, topK);97 98            const totalTopKProb = indexedProbs.reduce((a, b) => a + b.p, 0);99            let r = Math.random() * totalTopKProb;100            let nextToken = indexedProbs[0].i;101            102            for (let pObj of indexedProbs) {103                r -= pObj.p;104                if (r <= 0) {105                    nextToken = pObj.i;106                    break;107                }108            }109 110            if (nextToken === 50256) break; // EOS Token111 112            tokens.push(nextToken);113            114            // Dekodieren und senden115            const newText = enc.decode([nextToken]);116          117            if (clientConnected) {118                res.write(`data: ${JSON.stringify({ token: newText })}\n\n`);119            }120 121            // Event-Loop Pause122            await new Promise(r => setTimeout(r, 1));123        }124    } catch (err) {125        console.error("Fehler:", err);126        res.write(`data: ${JSON.stringify({ error: err.message })}\n\n`);127    } finally {128        res.end();129    }130});131 132app.get('/', (req, res) => res.send("SmaLLMPro Backend is Running"));133app.listen(7860, '0.0.0.0');134