331 lines
14 KiB
Python
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()
|