Files
llm-backend/localpilot/backends/transformers_backend.py
T

331 lines
14 KiB
Python

from __future__ import annotations
import asyncio
import gc
import queue
import threading
import time
from collections.abc import AsyncIterator
from pathlib import Path
from typing import Any
from ..config import ModelConfig, RuntimeConfig
from ..types import BackendError, ChatChunk, ChatParams
from .base import BaseBackend
class TransformersBackend(BaseBackend):
"""Load a Transformers/safetensors causal LM in the existing Conda LLM env."""
def __init__(self, model: ModelConfig, runtime: RuntimeConfig) -> None:
super().__init__(model, runtime)
self.tokenizer: Any = None
self.model_object: Any = None
self.device: Any = None
self._generation_lock = asyncio.Lock()
async def load(self) -> None:
if self.loaded:
return
if not self.model.path:
raise BackendError(f"Transformers 模型缺少 path: {self.model.id}")
path = Path(self.model.path)
if not path.exists():
raise BackendError(f"模型目录不存在: {path}")
try:
self.tokenizer, self.model_object, self.device = await asyncio.to_thread(self._load_sync, path)
except Exception as exc:
raise BackendError(f"Transformers 模型加载失败: {exc}") from exc
self.loaded = True
def _load_sync(self, path: Path) -> tuple[Any, Any, Any]:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
options = self.model.options
if torch.cuda.is_available():
try:
torch.backends.cuda.matmul.fp32_precision = "tf32"
torch.backends.cudnn.conv.fp32_precision = "tf32"
except AttributeError:
pass
model_dir = path if path.is_dir() else path.parent
dtype_name = str(options.get("torch_dtype", "float16")).lower()
dtype = {
"float16": torch.float16,
"fp16": torch.float16,
"bfloat16": torch.bfloat16,
"bf16": torch.bfloat16,
"float32": torch.float32,
"fp32": torch.float32,
}.get(dtype_name, torch.float16)
device_name = str(options.get("device", "cuda" if torch.cuda.is_available() else "cpu"))
if device_name == "cuda" and not torch.cuda.is_available():
device_name = "cpu"
tokenizer = AutoTokenizer.from_pretrained(model_dir, use_fast=True, local_files_only=True)
if tokenizer.pad_token_id is None and tokenizer.eos_token_id is not None:
tokenizer.pad_token = tokenizer.eos_token
load_kwargs: dict[str, Any] = {
"dtype": dtype,
"low_cpu_mem_usage": True,
}
if options.get("attn_implementation"):
load_kwargs["attn_implementation"] = options["attn_implementation"]
quantization = str(options.get("quantization", "none")).lower()
load_in_4bit = bool(options.get("load_in_4bit", False)) or quantization in {"4bit", "int4", "nf4"}
load_in_8bit = bool(options.get("load_in_8bit", False)) or quantization in {"8bit", "int8", "bnb8"}
dynamic_int8_cpu = bool(options.get("dynamic_int8_cpu", False)) or (
quantization in {"dynamic-int8", "int8-dynamic"} and device_name == "cpu"
)
if (load_in_4bit or load_in_8bit) and device_name.startswith("cuda"):
try:
from transformers import BitsAndBytesConfig
if load_in_4bit:
load_kwargs["quantization_config"] = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type=str(options.get("bnb_4bit_quant_type", "nf4")),
bnb_4bit_compute_dtype=dtype,
bnb_4bit_use_double_quant=bool(options.get("bnb_4bit_use_double_quant", True)),
)
else:
load_kwargs["quantization_config"] = BitsAndBytesConfig(load_in_8bit=True)
load_kwargs["device_map"] = options.get("device_map", "auto")
except Exception as exc:
raise BackendError(f"4-bit 量化加载失败,请检查 bitsandbytes/CUDA: {exc}") from exc
else:
try:
load_kwargs["device_map"] = options.get("device_map", "auto")
model_object = AutoModelForCausalLM.from_pretrained(model_dir, **load_kwargs)
except TypeError as first_exc:
if "dtype" not in str(first_exc):
raise
load_kwargs.pop("dtype", None)
load_kwargs["torch_dtype"] = dtype
model_object = AutoModelForCausalLM.from_pretrained(model_dir, **load_kwargs)
except (ImportError, ValueError) as first_exc:
load_kwargs.pop("device_map", None)
try:
model_object = AutoModelForCausalLM.from_pretrained(model_dir, **load_kwargs)
model_object.to(device_name)
except Exception as second_exc:
raise RuntimeError(f"自动设备映射失败: {first_exc}; 手动放置也失败: {second_exc}") from second_exc
if dynamic_int8_cpu and device_name == "cpu":
model_object = torch.quantization.quantize_dynamic(
model_object, {torch.nn.Linear}, dtype=torch.qint8
)
if "model_object" not in locals():
model_object = AutoModelForCausalLM.from_pretrained(model_dir, **load_kwargs)
adapters = options.get("adapters", [])
if isinstance(adapters, str):
adapters = [adapters]
for adapter in adapters:
try:
from peft import PeftModel
model_object = PeftModel.from_pretrained(model_object, adapter)
except ImportError as exc:
raise BackendError("配置了 ADAPTER,但当前环境缺少 peft") from exc
model_object.eval()
actual_device = next(model_object.parameters()).device
return tokenizer, model_object, actual_device
async def unload(self) -> None:
self.loaded = False
self.tokenizer = None
self.model_object = None
self.device = None
gc.collect()
try:
import torch
if torch.cuda.is_available():
torch.cuda.empty_cache()
except Exception:
pass
def _prepare_inputs_sync(self, messages: list[dict[str, Any]], params: ChatParams) -> dict[str, Any]:
prepared_messages: list[dict[str, Any]] = []
for message in messages:
item = dict(message)
content = item.get("content")
if isinstance(content, list):
item["content"] = "".join(
str(part.get("text", "")) if isinstance(part, dict) and part.get("type") == "text" else "[image]"
for part in content
)
elif not isinstance(content, str):
item["content"] = str(content)
prepared_messages.append(item)
kwargs: dict[str, Any] = {
"add_generation_prompt": True,
"tokenize": True,
"return_tensors": "pt",
"return_dict": True,
}
if self.model.template:
kwargs["chat_template"] = self.model.template
if params.enable_thinking is not None:
kwargs["enable_thinking"] = params.enable_thinking
try:
inputs = self.tokenizer.apply_chat_template(prepared_messages, **kwargs)
except TypeError:
kwargs.pop("enable_thinking", None)
inputs = self.tokenizer.apply_chat_template(prepared_messages, **kwargs)
return {key: value.to(self.device) if hasattr(value, "to") else value for key, value in inputs.items()}
def _generation_kwargs(self, params: ChatParams) -> dict[str, Any]:
kwargs: dict[str, Any] = {
"max_new_tokens": params.max_tokens,
"do_sample": params.temperature > 0,
"use_cache": True,
"repetition_penalty": params.repeat_penalty,
"pad_token_id": self.tokenizer.pad_token_id,
"eos_token_id": self.tokenizer.eos_token_id,
}
if params.temperature > 0:
kwargs["temperature"] = params.temperature
kwargs["top_k"] = params.top_k
kwargs["top_p"] = params.top_p
allowed_extra = {
"do_sample", "num_beams", "num_return_sequences", "typical_p", "epsilon_cutoff", "eta_cutoff",
"diversity_penalty", "length_penalty", "no_repeat_ngram_size", "max_length", "max_new_tokens",
"renormalize_logits", "remove_invalid_values", "bad_words_ids", "force_words_ids",
"repetition_penalty", "stop_strings",
}
kwargs.update({key: value for key, value in params.extra.items() if key in allowed_extra})
if params.seed is not None:
import torch
torch.manual_seed(params.seed)
return kwargs
def _clean_stop(self, text: str, params: ChatParams) -> str:
for stop in params.stop:
if stop in text:
text = text.split(stop, 1)[0]
return text
async def chat(
self,
messages: list[dict[str, Any]],
params: ChatParams,
stream: bool = False,
) -> AsyncIterator[ChatChunk]:
if not self.loaded:
await self.load()
async with self._generation_lock:
inputs = await asyncio.to_thread(self._prepare_inputs_sync, messages, params)
kwargs = self._generation_kwargs(params)
if not stream:
started = time.perf_counter()
output = await asyncio.to_thread(self.model_object.generate, **inputs, **kwargs)
prompt_len = int(inputs["input_ids"].shape[-1])
text = self.tokenizer.decode(output[0][prompt_len:], skip_special_tokens=True)
output_tokens = max(0, int(output.shape[-1]) - prompt_len)
duration = max(time.perf_counter() - started, 1e-9)
yield ChatChunk(
text=self._clean_stop(text, params),
done=True,
finish_reason="stop",
metadata={
"usage": {
"prompt_tokens": prompt_len,
"completion_tokens": output_tokens,
"total_tokens": prompt_len + output_tokens,
},
"timings": {
"total_seconds": duration,
"tokens_per_second": output_tokens / duration,
},
},
)
return
from transformers import TextIteratorStreamer
streamer = TextIteratorStreamer(self.tokenizer, skip_prompt=True, skip_special_tokens=True)
events: queue.Queue[tuple[str, Any]] = queue.Queue()
def worker() -> None:
try:
self.model_object.generate(**inputs, streamer=streamer, **kwargs)
except BaseException as exc:
events.put(("error", exc))
try:
streamer.end()
except Exception:
pass
def forward_stream() -> None:
try:
for item in streamer:
events.put(("text", item))
events.put(("done", None))
except BaseException as exc:
events.put(("error", exc))
generation_thread = threading.Thread(target=worker, name=f"localpilot-gen-{self.model.id}", daemon=True)
stream_thread = threading.Thread(target=forward_stream, name=f"localpilot-stream-{self.model.id}", daemon=True)
generation_thread.start()
stream_thread.start()
while True:
kind, value = await asyncio.to_thread(events.get)
if kind == "text":
text = self._clean_stop(str(value), params)
if text:
yield ChatChunk(text=text)
elif kind == "error":
raise BackendError(f"Transformers 推理失败: {value}") from value
elif kind == "done":
yield ChatChunk(done=True, finish_reason="stop")
return
def info(self) -> dict[str, Any]:
vram: dict[str, int] = {}
try:
import torch
if torch.cuda.is_available() and self.loaded:
vram = {
"vram_allocated_bytes": int(torch.cuda.memory_allocated()),
"vram_reserved_bytes": int(torch.cuda.memory_reserved()),
}
except Exception:
pass
return {
"id": self.model.id,
"kind": "transformers",
"path": self.model.path,
"loaded": self.loaded,
"device": str(self.device) if self.device is not None else None,
"load_in_4bit": bool(self.model.options.get("load_in_4bit", False)),
"quantization": self.model.options.get("quantization", "none"),
"dynamic_int8_cpu": bool(self.model.options.get("dynamic_int8_cpu", False)),
**vram,
}
async def embed(self, inputs: list[str]) -> list[list[float]]:
if not self.loaded:
await self.load()
return await asyncio.to_thread(self._embed_sync, inputs)
def _embed_sync(self, inputs: list[str]) -> list[list[float]]:
import torch
tokens = self.tokenizer(
inputs,
padding=True,
truncation=True,
return_tensors="pt",
)
tokens = {key: value.to(self.device) if hasattr(value, "to") else value for key, value in tokens.items()}
with torch.inference_mode():
outputs = self.model_object(**tokens, output_hidden_states=True, use_cache=False)
hidden = outputs.hidden_states[-1]
mask = tokens["attention_mask"].unsqueeze(-1).to(hidden.dtype)
pooled = (hidden * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1)
pooled = torch.nn.functional.normalize(pooled, p=2, dim=1)
return pooled.float().cpu().tolist()