atypique-api / core /multibackend_engine.py
RDS777's picture
feat: ajout de MultiBackendLLMEngine pour Qwen 7B local
3ff2a4f
Raw
History Blame Contribute Delete
3.57 kB
"""
MultiBackendLLMEngine – Moteur LLM local pour Qwen 7B (Apache 2.0)
Utilise transformers pour charger le modèle en 4-bit si disponible.
"""
import logging
import os
from typing import Optional
logger = logging.getLogger(__name__)
class MultiBackendLLMEngine:
"""Moteur LLM multi-backend (priorité au local Qwen)"""
def __init__(self, model_name: str = "Qwen/Qwen2.5-7B-Instruct"):
self.model_name = model_name
self.tokenizer = None
self.model = None
self.device = "cuda" if os.environ.get("CUDA_VISIBLE_DEVICES") else "cpu"
self._load_model()
def _load_model(self):
"""Charge le modèle et le tokenizer."""
try:
from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
import torch
logger.info(f"Chargement du modèle {self.model_name} sur {self.device}...")
# Tokenizer
self.tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
# Quantization 4-bit si GPU disponible
if self.device == "cuda":
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_use_double_quant=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16
)
self.model = AutoModelForCausalLM.from_pretrained(
self.model_name,
quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True
)
else:
self.model = AutoModelForCausalLM.from_pretrained(
self.model_name,
device_map="auto",
trust_remote_code=True
)
logger.info(f"✅ Modèle {self.model_name} chargé avec succès")
except Exception as e:
logger.error(f"Erreur lors du chargement du modèle: {e}")
self.model = None
self.tokenizer = None
raise
def generate(self, prompt: str, max_new_tokens: int = 512, temperature: float = 0.7) -> str:
"""Génère une réponse à partir d'un prompt."""
if not self.model or not self.tokenizer:
return "Modèle non chargé."
try:
inputs = self.tokenizer(prompt, return_tensors="pt", truncation=True, max_length=4096)
inputs = {k: v.to(self.model.device) for k, v in inputs.items()}
outputs = self.model.generate(
**inputs,
max_new_tokens=max_new_tokens,
temperature=temperature,
do_sample=True,
top_p=0.95,
pad_token_id=self.tokenizer.eos_token_id
)
response = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
# Enlever le prompt de la réponse
if response.startswith(prompt):
response = response[len(prompt):].strip()
return response
except Exception as e:
logger.error(f"Erreur lors de la génération: {e}")
return f"Erreur de génération: {e}"
def route(self, query: str) -> dict:
"""Point d'entrée compatible avec le MoERouter."""
response = self.generate(query)
return {"response": response, "source": "Qwen-7B-local"}