MORSATIPIK / core /multi_backend_engine.py
RDS777's picture
Update core/multi_backend_engine.py
9f8b820 verified
Raw
History Blame Contribute Delete
17.8 kB
# core/multi_backend_engine.py
import os
import logging
import asyncio
import aiohttp
import hashlib
import time
import traceback
from core.semantic_router import SemanticRouter
log = logging.getLogger("vortex.multibackend")
class LLMResponse:
__slots__ = ("content", "tokens", "cached")
def __init__(self, content: str, tokens: int = 0, cached: bool = False):
self.content = content
self.tokens = tokens
self.cached = cached
class BaseBackend:
def __init__(self):
self.failure_count = 0
self.last_failure_time = 0
self.quarantine_duration = 60
self.quarantined = False
def record_failure(self):
self.failure_count += 1
self.last_failure_time = time.time()
if self.failure_count >= 3:
self.quarantined = True
log.warning(f"[{self.__class__.__name__}] Mis en quarantaine après 3 échecs.")
else:
log.warning(f"[{self.__class__.__name__}] Échec #{self.failure_count}")
def record_success(self):
self.failure_count = 0
self.quarantined = False
self.last_failure_time = 0
async def is_available(self) -> bool:
return False
async def generate(self, model: str, system: str, user: str, max_tokens: int, temperature: float) -> str:
raise NotImplementedError
class OllamaBackend(BaseBackend):
def __init__(self, base_url="http://localhost"):
super().__init__()
self.base_url = base_url
self.ports = [11434, 11435, 11436]
self._active_port = None
self._session = None
async def _get_session(self):
if self._session is None or self._session.closed:
self._session = aiohttp.ClientSession()
return self._session
async def is_available(self) -> bool:
if self.quarantined:
if time.time() - self.last_failure_time > self.quarantine_duration:
self.quarantined = False
self.failure_count = 0
log.info(f"[OllamaBackend] Sortie de quarantaine.")
else:
return False
session = await self._get_session()
for port in self.ports:
try:
url = f"{self.base_url}:{port}/api/tags"
async with session.get(url, timeout=2) as resp:
if resp.status == 200:
self._active_port = port
return True
except Exception:
continue
return False
async def generate(self, model: str, system: str, user: str, max_tokens: int, temperature: float) -> str:
if not self._active_port and not await self.is_available():
self.record_failure()
return "[Ollama] Service indisponible."
session = await self._get_session()
url = f"{self.base_url}:{self._active_port}/api/chat"
payload = {
"model": model,
"messages": [
{"role": "system", "content": system or "Tu es VORTEX."},
{"role": "user", "content": user}
],
"options": {"num_predict": max_tokens, "temperature": temperature},
"stream": False
}
try:
async with session.post(url, json=payload, timeout=180) as resp:
if resp.status == 200:
data = await resp.json()
content = data.get("message", {}).get("content", "")
if content:
self.record_success()
return content
else:
self.record_failure()
return "[Ollama] Réponse vide."
self.record_failure()
return f"[Ollama] HTTP {resp.status}"
except Exception as e:
self.record_failure()
return f"[Ollama] Erreur: {e}"
class HFInferenceBackend(BaseBackend):
def __init__(self):
super().__init__()
self.token = os.environ.get("HF_TOKEN")
self._available = bool(self.token)
# Mise à jour du modèle pour l'API HF
self.default_model = "Qwen/Qwen2.5-Coder-7B-Instruct"
async def is_available(self) -> bool:
if self.quarantined:
if time.time() - self.last_failure_time > self.quarantine_duration:
self.quarantined = False
self.failure_count = 0
log.info(f"[HFInferenceBackend] Sortie de quarantaine.")
else:
return False
return self._available
async def generate(self, model: str, system: str, user: str, max_tokens: int, temperature: float) -> str:
if not self.token:
self.record_failure()
return "[HF] Token non configuré."
hf_model = self.default_model
try:
from huggingface_hub import InferenceClient
client = InferenceClient(token=self.token)
# Utilisation de chat_completion qui est plus robuste
messages = []
if system and system.strip():
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": user})
resp = client.chat_completion(
model=hf_model,
messages=messages,
max_tokens=max_tokens,
temperature=max(temperature, 0.01)
)
if resp and resp.choices and len(resp.choices) > 0:
content = resp.choices[0].message.content
if content and content.strip():
self.record_success()
return content.strip()
# Fallback sur text_generation si chat_completion échoue
try:
resp2 = client.text_generation(
user,
model=hf_model,
max_new_tokens=max_tokens,
temperature=max(temperature, 0.01),
do_sample=True
)
if resp2 and resp2.strip():
self.record_success()
return resp2.strip()
except Exception:
pass
self.record_failure()
return "[HF] Réponse vide."
except Exception as e:
self.record_failure()
return f"[HF] Erreur: {e}"
class LocalTransformersBackend(BaseBackend):
def __init__(self, error_tree=None):
super().__init__()
self._model = None
self._tokenizer = None
self._loaded = False
self._loading = False
self._torch_available = False
self._transformers_available = False
# 🔥 NOUVEAU MODÈLE : Qwen2.5-Coder-7B-Instruct
self.max_tokens_cap = int(os.environ.get("LOCAL_MAX_TOKENS_CAP", "512"))
self.timeout_s = float(os.environ.get("LOCAL_TIMEOUT_S", "45.0"))
self.model_name = "Qwen/Qwen2.5-Coder-7B-Instruct"
self.error_tree = error_tree
try:
import torch
self._torch_available = True
except ImportError:
log.warning("[Local] torch non installé.")
try:
import transformers
self._transformers_available = True
except ImportError:
log.warning("[Local] transformers non installé.")
def _build_prompt(self, system: str, user: str) -> str:
if self._tokenizer is None:
return user
messages = []
if system and system.strip():
messages.append({"role": "system", "content": system})
messages.append({"role": "user", "content": user})
try:
return self._tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
except Exception as e:
log.warning(f"[Local] apply_chat_template échoué: {e}, fallback user prompt")
return user
async def is_available(self) -> bool:
if not self._torch_available or not self._transformers_available:
return False
if self.quarantined:
if time.time() - self.last_failure_time > self.quarantine_duration:
self.quarantined = False
self.failure_count = 0
log.info(f"[LocalTransformersBackend] Sortie de quarantaine.")
else:
return False
if self._loaded:
return True
if self._loading:
return False
self._loading = True
try:
import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
log.info(f"[Local] Chargement de {self.model_name} en 4-bit NF4...")
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.float16,
)
self._tokenizer = AutoTokenizer.from_pretrained(
self.model_name,
trust_remote_code=True,
clean_up_tokenization_spaces=False
)
self._model = AutoModelForCausalLM.from_pretrained(
self.model_name,
quantization_config=bnb_config,
device_map="auto",
torch_dtype=torch.float16,
trust_remote_code=True
)
self._loaded = True
self._loading = False
log.info(f"[Local] {self.model_name} chargé avec succès.")
return True
except Exception as e:
self._loading = False
log.error(f"[Local] Échec du chargement: {e}")
log.error(traceback.format_exc())
if self.error_tree is not None:
try:
self.error_tree.add_error(str(e), f"LocalTransformers chargement")
except Exception:
pass
self.record_failure()
return False
async def generate(self, model: str, system: str, user: str, max_tokens: int, temperature: float) -> str:
if not self._loaded and not await self.is_available():
self.record_failure()
return "[Local] Modèle non chargé."
effective_max_tokens = min(max_tokens, self.max_tokens_cap)
prompt = self._build_prompt(system, user)
def _generate_sync():
import torch
t0 = time.perf_counter()
inputs = self._tokenizer(prompt, return_tensors="pt").to(self._model.device)
with torch.no_grad():
outputs = self._model.generate(
**inputs,
max_new_tokens=effective_max_tokens,
temperature=max(temperature, 0.01) if temperature > 0 else None,
do_sample=temperature > 0.0 and len(prompt) < 500,
pad_token_id=self._tokenizer.eos_token_id,
num_beams=1,
)
result = self._tokenizer.decode(
outputs[0][inputs.input_ids.shape[1]:],
skip_special_tokens=True
).strip()
t1 = time.perf_counter()
log.info(f"[Local] Génération en {t1-t0:.2f}s ({effective_max_tokens} max tokens)")
return result
try:
result = await asyncio.wait_for(
asyncio.to_thread(_generate_sync),
timeout=self.timeout_s
)
if result:
self.record_success()
return result
else:
self.record_failure()
if self.error_tree is not None:
try:
self.error_tree.add_error(
"Réponse vide du backend local",
f"Prompt: {prompt[:100]}"
)
except Exception:
pass
return "[Local] Réponse vide."
except asyncio.TimeoutError:
self.record_failure()
log.error(f"[Local] Génération trop lente (>{self.timeout_s}s).")
if self.error_tree is not None:
try:
self.error_tree.add_error(
f"Timeout local ({self.timeout_s}s)",
f"Prompt: {prompt[:100]}"
)
except Exception:
pass
return f"[Local] Timeout après {self.timeout_s}s."
except Exception as e:
self.record_failure()
log.error(f"[Local] Erreur génération: {type(e).__name__}: {e}")
log.error(traceback.format_exc())
if self.error_tree is not None:
try:
self.error_tree.add_error(str(e), f"LocalTransformers génération")
except Exception:
pass
return f"[Local] Erreur: {e}"
class MultiBackendLLMEngine:
def __init__(self, memory=None, error_tree=None):
self.memory = memory
self.error_tree = error_tree
self.router = SemanticRouter()
self._cache = {}
self.stats = {"total": 0, "backends": {}}
self.phi_model = os.environ.get("PHI_MODEL", "phi-4:latest")
self.gemma_model = os.environ.get("GEMMA_MODEL", "gemma4-moe:latest")
# Ordre des backends : LOCAL en PRIORITÉ, puis HF, puis Ollama
self.backends = []
try:
import torch
self.backends.append(LocalTransformersBackend(error_tree=self.error_tree))
log.info("[Multi] Backend local activé (Qwen2.5-Coder-7B-Instruct).")
except ImportError:
log.warning("[Multi] torch non trouvé, backend local désactivé.")
self.backends.append(HFInferenceBackend())
self.backends.append(OllamaBackend())
self._active_backend = None
self._backend_checked = False
async def _discover_backend(self):
if self._backend_checked and self._active_backend is not None:
still_ok = await self._active_backend.is_available()
if still_ok:
return self._active_backend
self._active_backend = None
self._backend_checked = False
for backend in self.backends:
ok = await backend.is_available()
if ok:
self._active_backend = backend
self._backend_checked = True
log.info(f"[Multi] Backend actif : {backend.__class__.__name__}")
return backend
if self.backends:
self._active_backend = self.backends[-1]
self._backend_checked = True
log.warning("[Multi] Aucun backend disponible, fallback sur le dernier backend.")
return self._active_backend
else:
log.error("[Multi] Aucun backend configuré.")
return None
async def call(self, agent=None, system="", user="", max_tokens=512, temperature=0.3,
use_cache=True, fast=False, profile=None, **kwargs) -> LLMResponse:
try:
action = self.router.route(user)
use_phi = action in ("reason", "act") or any(kw in user.lower() for kw in [
"code", "def ", "class ", "python", "function", "algorithm",
"optimize", "calculate", "math", "pytest", "assert", "import"
])
if fast:
use_phi = True
model = self.phi_model if use_phi else self.gemma_model
self.stats["total"] += 1
log.info(f"[Multi] Routing to {'Phi' if use_phi else 'Gemma'} (model: {model})")
key = hashlib.sha256(f"{model}|{system}|{user}|{max_tokens}|{temperature}".encode()).hexdigest()[:20]
if use_cache and key in self._cache:
return LLMResponse(self._cache[key], tokens=0, cached=True)
backend = await self._discover_backend()
if backend is None:
raise RuntimeError("Aucun backend valide disponible")
backend_name = backend.__class__.__name__
self.stats["backends"][backend_name] = self.stats["backends"].get(backend_name, 0) + 1
content = await backend.generate(model, system, user, max_tokens, temperature)
if isinstance(content, str) and content.startswith("[") and \
("Erreur" in content or "indisponible" in content or "non configuré" in content or "non chargé" in content):
self._active_backend = None
self._backend_checked = False
log.warning(f"[Multi] Backend {backend_name} a échoué. Redécouverte au prochain appel.")
else:
backend.record_success()
if use_cache and isinstance(content, str) and not content.startswith("["):
self._cache[key] = content
return LLMResponse(content, tokens=len(content.split()) if isinstance(content, str) else 0, cached=False)
except Exception as e:
log.error(f"[Multi] Exception in call: {e}")
log.error(traceback.format_exc())
if self.error_tree is not None:
try:
self.error_tree.add_error(str(e), f"MultiBackendLLMEngine call")
except Exception:
pass
return LLMResponse(f"[Erreur interne] {e}", tokens=0, cached=False)
def get_stats(self):
return self.stats
# Alias pour rétrocompatibilité
DualLLMEngine = MultiBackendLLMEngine