LeoHai123/SmaLLMPro
0
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 