import os
import json
import torch
import torch.nn as nn
from torch.nn import functional as F
from tokenizers import ByteLevelBPETokenizer
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from huggingface_hub import hf_hub_download
import urllib.request
import urllib.parse
import re
# --- 1. INITIALIZE FASTAPI API ---
app = FastAPI(title="Orbit SpaceStar Multi-Version Backend API", version="0.03_Super")
# Zorg dat je website (frontend) vlekkeloos mag praten met deze backend
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
print(f"🛰️ Orbit Multi-Engine active on device: {DEVICE.upper()}")
# --- 2. DYNAMIC ARCHITECTURE CONFIGURATION ---
MODEL_CONFIGS = {
"0.01": {"block_size": 256, "n_embd": 384, "n_head": 6, "n_layer": 8},
"0.02": {"block_size": 256, "n_embd": 384, "n_head": 6, "n_layer": 8}
}
MODEL_REPOS = {
"0.01": "Littendekitten/Orbit-SpaceStar-0.01",
"0.02": "Littendekitten/Orbit-SpaceStar-0.02"
}
loaded_models = {}
loaded_tokenizers = {}
# --- 3. TRANSFORMER DYNAMIC ARCHITECTURE ---
class Head(nn.Module):
def __init__(self, n_embd, head_size, block_size):
super().__init__()
self.key = nn.Linear(n_embd, head_size, bias=False)
self.query = nn.Linear(n_embd, head_size, bias=False)
self.value = nn.Linear(n_embd, head_size, bias=False)
self.register_buffer('tril', torch.tril(torch.ones(block_size, block_size)))
def forward(self, x):
B, T, C = x.shape
k, q = self.key(x), self.query(x)
wei = q @ k.transpose(-2, -1) * k.shape[-1]**-0.5
wei = wei.masked_fill(self.tril[:T, :T] == 0, float('-inf'))
wei = F.softmax(wei, dim=-1)
v = self.value(x)
return wei @ v
class MultiHeadAttention(nn.Module):
def __init__(self, n_embd, num_heads, head_size, block_size):
super().__init__()
self.heads = nn.ModuleList([Head(n_embd, head_size, block_size) for _ in range(num_heads)])
self.proj = nn.Sequential(nn.Linear(head_size * num_heads, n_embd))
def forward(self, x):
out = torch.cat([h(x) for h in self.heads], dim=-1)
return self.proj(out)
class FeedForward(nn.Module):
def __init__(self, n_embd):
super().__init__()
self.net = nn.Sequential(
nn.Linear(n_embd, 4 * n_embd),
nn.ReLU(),
nn.Linear(4 * n_embd, n_embd),
)
def forward(self, x):
return self.net(x)
class Block(nn.Module):
def __init__(self, n_embd, n_head, block_size):
super().__init__()
head_size = n_embd // n_head
self.sa = MultiHeadAttention(n_embd, n_head, head_size, block_size)
self.ffwd = FeedForward(n_embd)
self.ln1 = nn.LayerNorm(n_embd)
self.ln2 = nn.LayerNorm(n_embd)
def forward(self, x):
x = x + self.sa(self.ln1(x))
x = x + self.ffwd(self.ln2(x))
return x
class OrbitTransformer(nn.Module):
def __init__(self, vocab_size, config):
super().__init__()
self.block_size = config["block_size"]
self.token_embedding_table = nn.Embedding(vocab_size, config["n_embd"])
self.position_embedding_table = nn.Embedding(config["block_size"], config["n_embd"])
self.blocks = nn.Sequential(*[Block(config["n_embd"], n_head=config["n_head"], block_size=config["block_size"]) for _ in range(config["n_layer"])])
self.ln_f = nn.LayerNorm(config["n_embd"])
self.lm_head = nn.Linear(config["n_embd"], vocab_size)
def forward(self, idx):
B, T = idx.shape
tok_emb = self.token_embedding_table(idx)
pos_emb = self.position_embedding_table(torch.arange(T, device=idx.device))
x = tok_emb + pos_emb
x = self.blocks(x)
x = self.ln_f(x)
logits = self.lm_head(x)
return logits, None
def generate(self, idx, max_new_tokens, tokenizer, temperature=0.5, top_k=30):
"""Geüpgradede generatie met een keiharde herhalingsrem (Repetition Penalty)"""
end_token_id = tokenizer.token_to_id("<|end|>")
repetition_penalty = 1.3 # Strafpunten voor tokens die hij al gebruikt heeft (voorkomt loops!)
for _ in range(max_new_tokens):
idx_cond = idx[:, -self.block_size:]
logits, _ = self(idx_cond)
logits = logits[:, -1, :] / temperature
# REPETITION PENALTY: Als het model een token herhaalt, maken we de kans erop kleiner!
for token_id in set(idx[0].tolist()):
if logits[0, token_id] > 0:
logits[0, token_id] /= repetition_penalty
else:
logits[0, token_id] *= repetition_penalty
if top_k is not None:
v, _ = torch.topk(logits, min(top_k, logits.size(-1)))
logits[logits < v[:, [-1]]] = float('-inf')
probs = F.softmax(logits, dim=-1)
idx_next = torch.multinomial(probs, num_samples=1)
idx = torch.cat((idx, idx_next), dim=1)
# KEIHARDE STOP-TOKEN CHECK: Meteen kappen als hij klaar is!
if end_token_id is not None and idx_next.item() == end_token_id:
break
return idx
# --- 4. DYNAMIC MODEL LOADER ---
def get_model_and_tokenizer(version: str):
if version in loaded_models:
return loaded_models[version], loaded_tokenizers[version]
if version not in MODEL_REPOS:
raise HTTPException(status_code=400, detail=f"Versie {version} niet ondersteund.")
repo_id = MODEL_REPOS[version]
config = MODEL_CONFIGS.get(version, MODEL_CONFIGS["0.02"])
print(f"📥 Laden van versie {version} vanuit {repo_id}...")
try:
model_path = hf_hub_download(repo_id=repo_id, filename="best_model.pth")
meta_path = hf_hub_download(repo_id=repo_id, filename="OrbitTokenizer/vocab_meta.json")
vocab_path = hf_hub_download(repo_id=repo_id, filename="OrbitTokenizer/vocab.json")
merges_path = hf_hub_download(repo_id=repo_id, filename="OrbitTokenizer/merges.txt")
with open(meta_path, 'r', encoding='utf-8') as f:
vocab_meta = json.load(f)
vocab_size = vocab_meta["vocab_size"]
tokenizer = ByteLevelBPETokenizer(vocab_path, merges_path, add_prefix_space=True)
tokenizer.normalizer = None
model = OrbitTransformer(vocab_size, config)
model.load_state_dict(torch.load(model_path, map_location=DEVICE))
model.to(DEVICE)
model.eval()
loaded_models[version] = model
loaded_tokenizers[version] = tokenizer
return model, tokenizer
except Exception as e:
if version == "0.02":
print(f"⚠️ Proxy: Versie 0.02 faalt ({e}). Terugval naar 0.01!")
m01, t01 = get_model_and_tokenizer("0.01")
loaded_models["0.02"] = m01
loaded_tokenizers["0.02"] = t01
return m01, t01
else:
raise HTTPException(status_code=404, detail=f"Bestanden niet gevonden: {e}")
# --- 5. INTERNET SEARCH ENGINE (STRENGER AFGESTELD) ---
def fetch_internet_context(query: str) -> str:
"""Haalt info op via DuckDuckGo, maar ALLEEN bij duidelijke zoek-prompts."""
q_lower = query.lower().strip()
# Alleen zoeken als het echt een vraag is die begint met deze woorden
triggers = ["zoek naar", "wat is", "wie is", "hoe werkt", "search for", "what is", "who is", "tell me about"]
needs_internet = any(q_lower.startswith(t) for t in triggers)
# Extra handrem: als de prompt super kort is (zoals "HI" of "ramen"), NOOIT internet gebruiken!
if not needs_internet or len(q_lower) <= 4:
return ""
try:
clean_q = q_lower
for t in triggers:
if clean_q.startswith(t):
clean_q = clean_q[len(t):].strip()
break
search_term = urllib.parse.quote(clean_q if clean_q else query)
url = f"https://api.duckduckgo.com/?q={search_term}&format=json&no_html=1"
req = urllib.request.Request(url, headers={'User-Agent': 'Mozilla/5.0'})
with urllib.request.urlopen(req, timeout=3) as response:
data = json.loads(response.read().decode('utf-8'))
abstract = data.get("AbstractText", "")
if not abstract and data.get("RelatedTopics"):
for item in data["RelatedTopics"]:
if isinstance(item, dict) and "Text" in item:
abstract = item["Text"]
break
if abstract:
return abstract[:180].strip()
except Exception as e:
print(f"📡 API fout: {e}")
return ""
# --- 6. API ENDPOINTS ---
class ChatRequest(BaseModel):
prompt: str
version: str = "0.02"
temperature: float = 0.5 # Iets lager gezet voor meer stabiliteit en minder hallucinaties
max_tokens: int = 150 # Iets korter gezet zodat hij niet te lang door-rabbelt
@app.get("/")
def home():
return {"status": "Orbit API is Online", "device": DEVICE}
@app.post("/chat")
def chat(request: ChatRequest):
model, tokenizer = get_model_and_tokenizer(request.version)
# Internet check (negeert "HI")
internet_info = fetch_internet_context(request.prompt)
think_html_block = ""
# Bouw de nette chat-prompt
if internet_info:
think_html_block = f"Internet uplink succesvol.\nGezocht op: '{request.prompt}'\nGevonden info: {internet_info}\n"
formatted_prompt = f"<|user|>\n[Context: {internet_info}]\n{request.prompt}\n<|assistant|>\n"
else:
formatted_prompt = f"<|user|>\n{request.prompt}\n<|assistant|>\n"
# Tokenizen en sturen naar de GPU/CPU
context = torch.tensor([tokenizer.encode(formatted_prompt).ids], dtype=torch.long, device=DEVICE)
with torch.no_grad():
generated_tokens = model.generate(
context,
max_new_tokens=request.max_tokens,
tokenizer=tokenizer,
temperature=request.temperature,
top_k=30
)[0]
full_output = tokenizer.decode(generated_tokens.tolist(), skip_special_tokens=False)
# Sloop de prompt-tekst uit het uiteindelijke antwoord
generated_text = full_output.replace(formatted_prompt, "").replace("<|end|>", "").strip()
# --- DE ULTIEME ORBIT-SPACESTAR RETREATING & CLEANING FILTER ---
# Sloop alle rare achtergrond-gedachten eruit en forceer zijn naam!
# 1. Herken interne hersenspinsels en stop ze in jouw vette frontend blok!
if any(trigger in generated_text for trigger in ["Discuss ramen", "User greets", "User asks", "Provide the basic", "Analyze the"]):
parts = re.split(r'(Hello!|Hi!|I am|I think|To print|Based on)', generated_text, maxsplit=1)
if len(parts) >= 3:
internal_thought = parts[0].strip()
actual_speech = parts[1] + parts[2]
generated_text = f"Interne model-analyse:\n{internal_thought}\n{actual_speech}"
# 2. KEIHARDE KORRECTIE: Hij is NIET Orbit Ultra, hij is Orbit-SpaceStar!
generated_text = generated_text.replace("Orbit Ultra", "Orbit-SpaceStar")
# 3. Voorkom dubbele loops op je scherm (als hij twee keer dezelfde vraag aan zichzelf stelt, knippen we hem af)
if "What is the Python command" in generated_text:
sub_parts = generated_text.split("What is the Python command")
generated_text = sub_parts[0].strip()
# Plak het internet-denkblok (als dat er is) aan het antwoord vast
final_response = think_html_block + generated_text
return {
"version_used": request.version,
"response": final_response.strip()
}