用Transformers部署DiffusionGemma,提供OpenAI格式接口
Cimix
2026-06-12 16:53
1
干就完了,冲!
至于为什么不用vLLM…
因为是内网机器拉不下来镜像,而且我懒不想做太多编译
import asyncio
import json
import os
import re
import time
import uuid
from typing import Any
import torch
import uvicorn
from fastapi import FastAPI, HTTPException, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import StreamingResponse
from pydantic import BaseModel
from transformers import AutoProcessor, DiffusionGemmaForBlockDiffusion
os.environ["CUDA_VISIBLE_DEVICES"] = os.environ.get("CUDA_VISIBLE_DEVICES", "5")
model_path = os.environ.get(
"MODEL_PATH",
"/modelscope/google/diffusiongemma-26B-A4B-it",
)
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
CONFIG_DIR = os.environ.get("CONFIG_DIR", model_path)
CONFIG_FALLBACK_TO_BASE_DIR = os.environ.get("CONFIG_FALLBACK_TO_BASE_DIR", "1") == "1"
CONFIG_SOURCES: dict[str, str | None] = {}
SERVED_MODEL_NAME = os.environ.get("SERVED_MODEL_NAME", "gemma")
MODEL_OWNER = os.environ.get("MODEL_OWNER", "local")
def int_env(name: str, default: int) -> int:
try:
return int(os.environ.get(name, default))
except (TypeError, ValueError):
return default
MODEL_CREATED = int_env("MODEL_CREATED", 0)
def config_search_paths() -> list[str]:
paths = [CONFIG_DIR]
if CONFIG_FALLBACK_TO_BASE_DIR:
paths.append(BASE_DIR)
unique_paths = []
for path in paths:
normalized_path = os.path.abspath(os.path.expanduser(path))
if normalized_path not in unique_paths:
unique_paths.append(normalized_path)
return unique_paths
def load_json_config(filename: str) -> dict[str, Any]:
for base_path in config_search_paths():
candidate = os.path.join(base_path, filename)
if not os.path.isfile(candidate):
continue
with open(candidate, "r", encoding="utf-8") as f:
CONFIG_SOURCES[filename] = candidate
return json.load(f)
CONFIG_SOURCES[filename] = None
return {}
TOKENIZER_CONFIG = load_json_config("tokenizer_config.json")
GENERATION_CONFIG = load_json_config("generation_config.json")
MODEL_CONFIG = load_json_config("config.json")
for config_name, config_source in CONFIG_SOURCES.items():
print(f"[config] {config_name}: {config_source or 'NOT FOUND'}")
print(f"[config] model_path: {model_path}")
print(f"[config] config_dir: {os.path.abspath(os.path.expanduser(CONFIG_DIR))}")
print(f"[config] fallback_to_base_dir: {CONFIG_FALLBACK_TO_BASE_DIR}")
def int_config(config: dict[str, Any], name: str, default: int) -> int:
value = config.get(name, default)
try:
return int(value)
except (TypeError, ValueError):
return default
processor = AutoProcessor.from_pretrained(
model_path,
trust_remote_code=True,
local_files_only=True,
)
model = DiffusionGemmaForBlockDiffusion.from_pretrained(
model_path,
device_map="auto",
trust_remote_code=True,
torch_dtype=torch.bfloat16,
local_files_only=True,
)
print(f"DiffusionGemma BF16 loaded | GPU:5 | Dtype: {model.dtype}")
app = FastAPI()
app.add_middleware(
CORSMiddleware,
allow_origins=os.environ.get("CORS_ALLOW_ORIGINS", "*").split(","),
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
MODEL_SPECIFIC_TOKENS = TOKENIZER_CONFIG.get("model_specific_special_tokens", {})
CONFIG_MAX_NEW_TOKENS = int_config(GENERATION_CONFIG, "max_new_tokens", 256)
DEFAULT_MAX_COMPLETION_TOKENS = int_env(
"DEFAULT_MAX_COMPLETION_TOKENS",
CONFIG_MAX_NEW_TOKENS,
)
DEFAULT_NUM_INFERENCE_STEPS = int_config(
GENERATION_CONFIG,
"max_denoising_steps",
48,
)
MODEL_CANVAS_LENGTH = int_config(MODEL_CONFIG, "canvas_length", 0)
MODEL_MAX_POSITION_EMBEDDINGS = int_config(
MODEL_CONFIG.get("text_config", {}),
"max_position_embeddings",
int_config(MODEL_CONFIG, "max_position_embeddings", 0),
)
def configured_token(name: str, default: str) -> str:
value = TOKENIZER_CONFIG.get(name) or MODEL_SPECIFIC_TOKENS.get(name)
return value if isinstance(value, str) else default
BOS_TOKEN = configured_token("bos_token", "<bos>")
EOS_TOKEN = configured_token("eos_token", "<eos>")
PAD_TOKEN = configured_token("pad_token", "<pad>")
SOT_TOKEN = configured_token("sot_token", "<|turn>")
EOT_TOKEN = configured_token("eot_token", "<turn|>")
SOC_TOKEN = configured_token("soc_token", "<|channel>")
EOC_TOKEN = configured_token("eoc_token", "<channel|>")
STC_TOKEN = configured_token("stc_token", "<|tool_call>")
ETC_TOKEN = configured_token("etc_token", "<tool_call|>")
STD_TOKEN = configured_token("std_token", "<|tool>")
ETD_TOKEN = configured_token("etd_token", "<tool|>")
STR_TOKEN = configured_token("str_token", "<|tool_response>")
ETR_TOKEN = configured_token("etr_token", "<tool_response|>")
THINK_TOKEN = configured_token("think_token", "<|think|>")
ESCAPE_TOKEN = configured_token("escape_token", '<|"|>')
MODEL_TURN_PREFIX = f"{SOT_TOKEN}model\n"
ASSISTANT_TURN_PREFIX = f"{SOT_TOKEN}assistant\n"
THOUGHT_CHANNEL_PREFIX = f"{SOC_TOKEN}thought\n"
TOOL_CALL_BLOCK_RE = re.compile(
rf"{re.escape(STC_TOKEN)}(.*?){re.escape(ETC_TOKEN)}",
re.DOTALL,
)
GEMMA_TOOL_CALL_RE = re.compile(
r"call:(?P<name>\w+)\s*(?P<arguments>\{.*\})?",
re.DOTALL,
)
def collect_control_tokens() -> tuple[str, ...]:
tokens = {
BOS_TOKEN,
EOS_TOKEN,
PAD_TOKEN,
SOT_TOKEN,
EOT_TOKEN,
SOC_TOKEN,
EOC_TOKEN,
STC_TOKEN,
ETC_TOKEN,
STD_TOKEN,
ETD_TOKEN,
STR_TOKEN,
ETR_TOKEN,
THINK_TOKEN,
ESCAPE_TOKEN,
MODEL_TURN_PREFIX.rstrip("\n"),
ASSISTANT_TURN_PREFIX.rstrip("\n"),
THOUGHT_CHANNEL_PREFIX.rstrip("\n"),
}
def add_tokens(value: Any) -> None:
if isinstance(value, str) and value.startswith("<"):
tokens.add(value)
elif isinstance(value, list):
for item in value:
add_tokens(item)
elif isinstance(value, dict):
for item in value.values():
add_tokens(item)
for key, value in TOKENIZER_CONFIG.items():
if key.endswith("_token") or key in {
"extra_special_tokens",
"model_specific_special_tokens",
}:
add_tokens(value)
return tuple(sorted(tokens, key=len, reverse=True))
CONTROL_TOKENS = collect_control_tokens()
USE_EXPLICIT_GENERATION_TOKENS = (
os.environ.get("USE_EXPLICIT_GENERATION_TOKENS", "0") == "1"
)
DEBUG_DECODE = os.environ.get("DEBUG_DECODE", "0") == "1"
DEBUG_LENGTH = os.environ.get("DEBUG_LENGTH", "0") == "1"
def served_model_ids() -> list[str]:
aliases = [
alias.strip()
for alias in os.environ.get("MODEL_ALIASES", "").split(",")
if alias.strip()
]
path_model_name = os.path.basename(os.path.normpath(model_path))
candidates = [SERVED_MODEL_NAME, path_model_name, *aliases]
model_ids = []
for model_id in candidates:
if model_id and model_id not in model_ids:
model_ids.append(model_id)
return model_ids or ["gemma"]
MODEL_IDS = served_model_ids()
def model_card(model_id: str) -> dict[str, Any]:
return {
"id": model_id,
"object": "model",
"created": MODEL_CREATED,
"owned_by": MODEL_OWNER,
"model_path": model_path,
"config_dir": os.path.abspath(os.path.expanduser(CONFIG_DIR)),
"config_sources": CONFIG_SOURCES,
"max_context_length": MODEL_MAX_POSITION_EMBEDDINGS or None,
"canvas_length": MODEL_CANVAS_LENGTH or None,
"default_max_completion_tokens": DEFAULT_MAX_COMPLETION_TOKENS,
"config_max_new_tokens": CONFIG_MAX_NEW_TOKENS,
}
@app.get("/")
@app.get("/v1")
async def root():
return {
"object": "service",
"status": "ok",
"model": MODEL_IDS[0],
"model_path": model_path,
"config_dir": os.path.abspath(os.path.expanduser(CONFIG_DIR)),
"config_sources": CONFIG_SOURCES,
"max_context_length": MODEL_MAX_POSITION_EMBEDDINGS or None,
"canvas_length": MODEL_CANVAS_LENGTH or None,
"default_max_completion_tokens": DEFAULT_MAX_COMPLETION_TOKENS,
"endpoints": [
"/v1/models",
"/v1/chat/completions",
"/health",
],
}
@app.get("/health")
@app.get("/v1/health")
async def health():
return {
"status": "ok",
"model_loaded": True,
"model": MODEL_IDS[0],
"model_path": model_path,
"config_dir": os.path.abspath(os.path.expanduser(CONFIG_DIR)),
"config_sources": CONFIG_SOURCES,
"dtype": str(getattr(model, "dtype", "")),
"max_context_length": MODEL_MAX_POSITION_EMBEDDINGS or None,
"canvas_length": MODEL_CANVAS_LENGTH or None,
"default_max_completion_tokens": DEFAULT_MAX_COMPLETION_TOKENS,
"config_max_new_tokens": CONFIG_MAX_NEW_TOKENS,
}
@app.get("/models")
@app.get("/v1/models")
async def list_models():
return {
"object": "list",
"data": [model_card(model_id) for model_id in MODEL_IDS],
}
@app.get("/models/{model_id:path}")
@app.get("/v1/models/{model_id:path}")
async def retrieve_model(model_id: str):
if model_id not in MODEL_IDS:
raise HTTPException(status_code=404, detail=f"Model {model_id!r} not found")
return model_card(model_id)
def generation_token_kwargs() -> dict[str, Any]:
if not USE_EXPLICIT_GENERATION_TOKENS:
return {}
pad_token_id = GENERATION_CONFIG.get(
"pad_token_id",
getattr(processor.tokenizer, "pad_token_id", None),
)
eos_token_id = GENERATION_CONFIG.get(
"eos_token_id",
getattr(processor.tokenizer, "eos_token_id", None),
)
kwargs = {}
if pad_token_id is not None:
kwargs["pad_token_id"] = pad_token_id
if eos_token_id is not None:
kwargs["eos_token_id"] = eos_token_id
return kwargs
def coerce_decode_result_to_text(value) -> str:
"""DiffusionGemma's processor.decode can return either str or list."""
if value is None:
return ""
if isinstance(value, str):
return value
if isinstance(value, (list, tuple)):
return "".join(coerce_decode_result_to_text(item) for item in value)
return str(value)
def decode_tokens(tokens, *, skip_special_tokens: bool = False) -> str:
decoded = processor.decode(tokens, skip_special_tokens=skip_special_tokens)
return coerce_decode_result_to_text(decoded)
def clean_generated_text(text: str) -> str:
if MODEL_TURN_PREFIX.rstrip("\n") in text:
text = text.split(MODEL_TURN_PREFIX.rstrip("\n"))[-1]
if ASSISTANT_TURN_PREFIX.rstrip("\n") in text:
text = text.split(ASSISTANT_TURN_PREFIX.rstrip("\n"))[-1]
for token in CONTROL_TOKENS:
text = text.replace(token, "")
return text.strip()
def strip_padding_tokens(text: str) -> str:
return (
text.replace("\r\n", "\n")
.replace(BOS_TOKEN, "")
.replace(EOS_TOKEN, "")
.replace(PAD_TOKEN, "")
.strip()
)
def isolate_last_model_turn(text: str) -> str:
text = strip_padding_tokens(text)
if MODEL_TURN_PREFIX in text:
text = text.rsplit(MODEL_TURN_PREFIX, 1)[-1]
elif ASSISTANT_TURN_PREFIX in text:
text = text.rsplit(ASSISTANT_TURN_PREFIX, 1)[-1]
elif MODEL_TURN_PREFIX.rstrip("\n") in text:
text = text.rsplit(MODEL_TURN_PREFIX.rstrip("\n"), 1)[-1].lstrip()
elif ASSISTANT_TURN_PREFIX.rstrip("\n") in text:
text = text.rsplit(ASSISTANT_TURN_PREFIX.rstrip("\n"), 1)[-1].lstrip()
for stop_token in (EOT_TOKEN, STR_TOKEN, ETR_TOKEN):
if stop_token in text:
text = text.split(stop_token, 1)[0]
return text.strip()
class GemmaToolArgumentParser:
def __init__(self, text: str):
self.text = text.strip()
self.pos = 0
def current(self) -> str:
if self.pos >= len(self.text):
return ""
return self.text[self.pos]
def startswith(self, value: str) -> bool:
return self.text.startswith(value, self.pos)
def skip_whitespace(self) -> None:
while self.current() and self.current().isspace():
self.pos += 1
def parse(self) -> Any:
self.skip_whitespace()
return self.parse_value()
def parse_value(self) -> Any:
self.skip_whitespace()
if self.startswith(ESCAPE_TOKEN):
return self.parse_escaped_string()
if self.current() == "{":
return self.parse_object()
if self.current() == "[":
return self.parse_array()
return self.parse_bare_value()
def parse_escaped_string(self) -> str:
self.pos += len(ESCAPE_TOKEN)
end = self.text.find(ESCAPE_TOKEN, self.pos)
if end < 0:
value = self.text[self.pos :]
self.pos = len(self.text)
return value
value = self.text[self.pos : end]
self.pos = end + len(ESCAPE_TOKEN)
return value
def parse_key(self) -> str:
self.skip_whitespace()
if self.startswith(ESCAPE_TOKEN):
return self.parse_escaped_string()
start = self.pos
while self.current() and self.current() not in ":{}[],":
self.pos += 1
return self.text[start : self.pos].strip().strip('"')
def parse_object(self) -> dict[str, Any]:
result = {}
self.pos += 1
while self.current():
self.skip_whitespace()
if self.current() == "}":
self.pos += 1
break
key = self.parse_key()
self.skip_whitespace()
if self.current() != ":":
break
self.pos += 1
result[key] = self.parse_value()
self.skip_whitespace()
if self.current() == ",":
self.pos += 1
continue
if self.current() == "}":
self.pos += 1
break
return result
def parse_array(self) -> list[Any]:
result = []
self.pos += 1
while self.current():
self.skip_whitespace()
if self.current() == "]":
self.pos += 1
break
result.append(self.parse_value())
self.skip_whitespace()
if self.current() == ",":
self.pos += 1
continue
if self.current() == "]":
self.pos += 1
break
return result
def parse_bare_value(self) -> Any:
start = self.pos
while self.current() and self.current() not in ",]}":
self.pos += 1
raw = self.text[start : self.pos].strip()
lower = raw.lower()
if lower == "true":
return True
if lower == "false":
return False
if lower in {"null", "none"}:
return None
if re.fullmatch(r"[-+]?\d+", raw):
try:
return int(raw)
except ValueError:
pass
if re.fullmatch(r"[-+]?(\d+\.\d*|\d*\.\d+)([eE][-+]?\d+)?", raw):
try:
return float(raw)
except ValueError:
pass
return raw.strip('"')
def parse_gemma_arguments(arguments: str) -> dict[str, Any]:
text = arguments.strip()
if not text:
return {}
try:
parsed = json.loads(text.replace(ESCAPE_TOKEN, '"'))
if isinstance(parsed, dict):
return parsed
except json.JSONDecodeError:
pass
try:
parsed = GemmaToolArgumentParser(text).parse()
if isinstance(parsed, dict):
return parsed
except Exception:
pass
return {"_raw": text}
def parse_arguments_to_object(arguments: Any) -> dict[str, Any]:
if isinstance(arguments, dict):
return arguments
if arguments is None:
return {}
if isinstance(arguments, str):
try:
parsed = json.loads(arguments)
if isinstance(parsed, dict):
return parsed
except json.JSONDecodeError:
pass
return parse_gemma_arguments(arguments)
return {"value": arguments}
def normalize_internal_tool_call(tool_call: Any) -> dict[str, Any] | None:
if not isinstance(tool_call, dict):
return None
function = tool_call.get("function")
if not isinstance(function, dict):
function = tool_call
name = function.get("name")
if not name:
return None
normalized = {
"type": "function",
"function": {
"name": str(name),
"arguments": parse_arguments_to_object(function.get("arguments", {})),
},
}
if tool_call.get("id"):
normalized["id"] = str(tool_call["id"])
return normalized
def normalize_internal_tool_calls(tool_calls: Any) -> list[dict[str, Any]]:
if not isinstance(tool_calls, list):
return []
normalized = []
for tool_call in tool_calls:
normalized_call = normalize_internal_tool_call(tool_call)
if normalized_call:
normalized.append(normalized_call)
return normalized
def parse_tool_calls_from_text(text: str) -> list[dict[str, Any]]:
tool_calls = []
for match in TOOL_CALL_BLOCK_RE.finditer(text):
body = match.group(1).strip()
call_match = GEMMA_TOOL_CALL_RE.search(body)
if not call_match:
continue
name = call_match.group("name")
arguments = call_match.group("arguments") or "{}"
tool_calls.append(
{
"type": "function",
"function": {
"name": name,
"arguments": parse_gemma_arguments(arguments),
},
}
)
return tool_calls
def split_response_parts(decoded_text: str) -> tuple[str, str, list[dict[str, Any]]]:
text = isolate_last_model_turn(decoded_text)
thinking = ""
content = text
start = content.find(THOUGHT_CHANNEL_PREFIX)
if start >= 0:
before_thinking = content[:start].strip()
after_start = content[start + len(THOUGHT_CHANNEL_PREFIX):]
if EOC_TOKEN in after_start:
thinking, content = after_start.split(EOC_TOKEN, 1)
if before_thinking:
content = before_thinking + "\n" + content
else:
thinking = after_start
content = before_thinking
tool_calls = parse_tool_calls_from_text(content)
content = TOOL_CALL_BLOCK_RE.sub("", content)
thinking = clean_generated_text(thinking)
content = clean_generated_text(content)
return thinking, content, tool_calls
def parse_response_text(decoded_text: str) -> tuple[str, str, list[dict[str, Any]]]:
fallback_thinking, fallback_content, fallback_tool_calls = split_response_parts(
decoded_text
)
parse_response = getattr(processor, "parse_response", None)
if callable(parse_response):
try:
parsed = parse_response(decoded_text)
except Exception:
parsed = None
if isinstance(parsed, dict):
thinking = clean_generated_text(
coerce_decode_result_to_text(parsed.get("thinking", ""))
)
content = TOOL_CALL_BLOCK_RE.sub(
"",
coerce_decode_result_to_text(parsed.get("content", "")),
)
content = clean_generated_text(content)
tool_calls = normalize_internal_tool_calls(parsed.get("tool_calls", []))
if not thinking:
thinking = fallback_thinking
if not content:
content = fallback_content
if not tool_calls:
tool_calls = fallback_tool_calls
parsed_visible = select_visible_content(thinking, content)
fallback_visible = select_visible_content(
fallback_thinking,
fallback_content,
)
if fallback_visible and not parsed_visible and not tool_calls:
return fallback_thinking, fallback_content, fallback_tool_calls
if thinking or content or tool_calls:
return thinking, content, tool_calls
return fallback_thinking, fallback_content, fallback_tool_calls
ROLE_ARTIFACTS = {"system", "user", "assistant", "model"}
ASSISTANT_ROLE_ARTIFACTS = {"assistant", "model"}
PROMPT_ROLE_ARTIFACTS = {"system", "user"}
def is_role_artifact(text: str) -> bool:
return text.strip().lower() in ROLE_ARTIFACTS
def split_leading_role_line(text: str) -> tuple[str | None, str]:
lines = text.strip().splitlines()
if not lines:
return None, ""
first_line = lines[0].strip().lower()
if first_line not in ROLE_ARTIFACTS:
return None, text.strip()
return first_line, "\n".join(lines[1:]).strip()
def normalize_visible_role_content(content: str) -> str:
content = content.strip()
while content:
role, rest = split_leading_role_line(content)
if role is None:
return content
if role in ASSISTANT_ROLE_ARTIFACTS:
content = rest
continue
if role in PROMPT_ROLE_ARTIFACTS:
return ""
return ""
def recover_content_from_thinking(thinking: str) -> str:
"""Recover only a quoted final answer, not the whole reasoning trace."""
candidates = []
for pattern in (r"“([^”]{2,500})”", r'"([^"\n]{2,500})"'):
candidates.extend(match.strip() for match in re.findall(pattern, thinking))
for line in reversed(thinking.splitlines()):
line = line.strip()
if not line:
continue
for prefix in ("最终答案:", "最终回答:", "答案:", "回复:", "Final answer:", "Answer:"):
if line.startswith(prefix):
candidates.append(line[len(prefix):].strip())
for candidate in reversed(candidates):
candidate = candidate.strip()
if not candidate or is_role_artifact(candidate):
continue
if "?" in candidate or "?" in candidate:
continue
if candidate:
return candidate
return ""
def format_openai_content(content: str) -> str:
return content.strip()
def select_visible_content(thinking: str, content: str) -> str:
content = normalize_visible_role_content(format_openai_content(content))
if content and not is_role_artifact(content):
return content
recovered = recover_content_from_thinking(thinking)
if recovered:
return recovered
return "" if is_role_artifact(content) else content
def extract_model_response(
outputs,
inputs,
) -> tuple[str, str, list[dict[str, Any]], int]:
prompt_len = inputs["input_ids"].shape[-1]
output_tokens = outputs[0]
output_len = len(output_tokens)
generated_len = max(0, output_len - prompt_len)
full_decoded_text = decode_tokens(output_tokens, skip_special_tokens=False)
decoded_candidates = [full_decoded_text]
if output_len > prompt_len:
sliced_decoded_text = decode_tokens(
output_tokens[prompt_len:],
skip_special_tokens=False,
)
if sliced_decoded_text != full_decoded_text:
decoded_candidates.append(sliced_decoded_text)
fallback_response = ("", "", [], generated_len)
for index, decoded_text in enumerate(decoded_candidates):
if DEBUG_DECODE:
print(f"[decode:{index}] {decoded_text[:2000]!r}")
thinking, content, tool_calls = parse_response_text(decoded_text)
if thinking or content or tool_calls:
fallback_response = (thinking, content, tool_calls, generated_len)
if tool_calls or select_visible_content(thinking, content):
return thinking, content, tool_calls, generated_len
return fallback_response
def request_field_was_set(req: "ChatRequest", field_name: str) -> bool:
fields_set = getattr(req, "model_fields_set", None)
if fields_set is None:
fields_set = getattr(req, "__fields_set__", set())
return field_name in fields_set
def resolve_requested_max_tokens(req: "ChatRequest") -> tuple[int, str]:
if request_field_was_set(req, "max_completion_tokens") and (
req.max_completion_tokens is not None
):
return int(req.max_completion_tokens), "max_completion_tokens"
if request_field_was_set(req, "max_tokens") and req.max_tokens is not None:
return int(req.max_tokens), "max_tokens"
if req.max_completion_tokens is not None:
return int(req.max_completion_tokens), "max_completion_tokens_default"
return DEFAULT_MAX_COMPLETION_TOKENS, "max_completion_tokens_default"
def get_request_max_tokens(req: "ChatRequest", prompt_tokens: int) -> int:
value, _source = resolve_requested_max_tokens(req)
requested_max_tokens = max(1, int(value))
if MODEL_MAX_POSITION_EMBEDDINGS <= 0:
return requested_max_tokens
available_tokens = MODEL_MAX_POSITION_EMBEDDINGS - prompt_tokens
if available_tokens < 1:
raise HTTPException(
status_code=400,
detail=(
"Input is too long for this model context window: "
f"prompt_tokens={prompt_tokens}, "
f"max_position_embeddings={MODEL_MAX_POSITION_EMBEDDINGS}"
),
)
return min(requested_max_tokens, available_tokens)
def build_generation_kwargs(
inputs: dict[str, Any],
req: "ChatRequest",
) -> dict[str, Any]:
kwargs = {key: value for key, value in inputs.items()}
temperature = req.temperature if req.temperature is not None else 0
prompt_tokens = inputs["input_ids"].shape[-1]
requested_max_tokens, max_tokens_source = resolve_requested_max_tokens(req)
effective_max_tokens = get_request_max_tokens(req, prompt_tokens)
if DEBUG_LENGTH:
available_tokens = (
MODEL_MAX_POSITION_EMBEDDINGS - prompt_tokens
if MODEL_MAX_POSITION_EMBEDDINGS > 0
else None
)
print(
"[length] "
f"source={max_tokens_source} "
f"requested={requested_max_tokens} "
f"effective={effective_max_tokens} "
f"prompt_tokens={prompt_tokens} "
f"available_tokens={available_tokens} "
f"context={MODEL_MAX_POSITION_EMBEDDINGS or None} "
f"canvas_length={MODEL_CANVAS_LENGTH or None}"
)
kwargs.update(
{
"max_new_tokens": effective_max_tokens,
"num_inference_steps": req.num_inference_steps,
"do_sample": temperature > 0,
}
)
if temperature > 0:
kwargs["temperature"] = temperature
kwargs.update(generation_token_kwargs())
return {key: value for key, value in kwargs.items() if value is not None}
def tool_function_name(tool: Any) -> str | None:
if not isinstance(tool, dict):
return None
function = tool.get("function")
if isinstance(function, dict):
name = function.get("name")
else:
name = tool.get("name")
return str(name) if name else None
def normalize_tools_for_processor(tools: Any) -> list[dict[str, Any]] | None:
if not tools:
return None
normalized = []
for tool in tools:
if not isinstance(tool, dict):
continue
function = tool.get("function")
if isinstance(function, dict):
function = dict(function)
function.setdefault("parameters", {"type": "object", "properties": {}})
normalized.append({"type": "function", "function": function})
elif tool.get("name"):
function = dict(tool)
function.setdefault("parameters", {"type": "object", "properties": {}})
normalized.append({"type": "function", "function": function})
return normalized or None
def select_tools_for_request(
tools: Any,
tool_choice: Any,
) -> tuple[list[dict[str, Any]] | None, str | None]:
normalized_tools = normalize_tools_for_processor(tools)
if not normalized_tools:
return None, None
if tool_choice == "none":
return None, None
if tool_choice in (None, "auto"):
return normalized_tools, None
if tool_choice == "required":
return (
normalized_tools,
"You must call one of the available tools when answering this turn.",
)
if isinstance(tool_choice, dict):
function = tool_choice.get("function")
name = function.get("name") if isinstance(function, dict) else None
if not name:
return normalized_tools, None
selected = [
tool for tool in normalized_tools if tool_function_name(tool) == name
]
if not selected:
raise HTTPException(
status_code=400,
detail=f"tool_choice function {name!r} is not present in tools",
)
return selected, f"You must call the function named {name}."
return normalized_tools, None
def parse_tool_response_value(value: Any) -> Any:
if not isinstance(value, str):
return value
try:
return json.loads(value)
except json.JSONDecodeError:
return value
def normalize_tool_response(
message: dict[str, Any],
tool_call_names: dict[str, str],
) -> dict[str, Any]:
tool_call_id = message.get("tool_call_id")
name = message.get("name")
if not name and tool_call_id:
name = tool_call_names.get(str(tool_call_id))
response = {"response": parse_tool_response_value(message.get("content", ""))}
if name:
response["name"] = str(name)
return response
def fallback_tool_response_message(tool_response: dict[str, Any]) -> dict[str, str]:
name = tool_response.get("name", "tool")
response = tool_response.get("response", "")
if not isinstance(response, str):
response = json.dumps(response, ensure_ascii=False)
return {"role": "user", "content": f"Tool response from {name}: {response}"}
def normalize_messages_for_processor(
messages: list[dict[str, Any]],
) -> list[dict[str, Any]]:
normalized_messages = []
tool_call_names = {}
last_assistant_message = None
for raw_message in messages:
if not isinstance(raw_message, dict):
normalized_messages.append(raw_message)
last_assistant_message = None
continue
role = raw_message.get("role")
if role == "tool":
tool_response = normalize_tool_response(raw_message, tool_call_names)
if last_assistant_message is not None:
last_assistant_message.setdefault("tool_responses", []).append(
tool_response
)
else:
normalized_messages.append(fallback_tool_response_message(tool_response))
last_assistant_message = None
continue
message = dict(raw_message)
if role == "assistant":
original_tool_calls = raw_message.get("tool_calls") or []
normalized_tool_calls = []
for original_tool_call in original_tool_calls:
normalized_tool_call = normalize_internal_tool_call(original_tool_call)
if not normalized_tool_call:
continue
normalized_tool_calls.append(normalized_tool_call)
if isinstance(original_tool_call, dict) and original_tool_call.get("id"):
tool_call_names[str(original_tool_call["id"])] = (
normalized_tool_call["function"]["name"]
)
if normalized_tool_calls:
message["tool_calls"] = normalized_tool_calls
if message.get("content") is None:
message["content"] = ""
if message.get("tool_responses"):
message["tool_responses"] = [
normalize_tool_response(response, tool_call_names)
for response in message["tool_responses"]
if isinstance(response, dict)
]
last_assistant_message = message
else:
last_assistant_message = None
normalized_messages.append(message)
return normalized_messages
def prepend_system_message(
messages: list[dict[str, Any]],
content: str,
) -> list[dict[str, Any]]:
if not content:
return messages
if messages and isinstance(messages[0], dict) and messages[0].get("role") == "system":
first = dict(messages[0])
original_content = first.get("content") or ""
first["content"] = (
f"{original_content}\n\n{content}" if original_content else content
)
return [first, *messages[1:]]
return [{"role": "system", "content": content}, *messages]
def tool_prompt_fallback(tools: list[dict[str, Any]]) -> str:
serialized_tools = json.dumps(tools, ensure_ascii=False)
return (
"Available tools are provided as JSON schemas below. If a tool is needed, "
f"respond with one or more tool call blocks in this exact form: "
f"{STC_TOKEN}call:function_name{{key:{ESCAPE_TOKEN}value{ESCAPE_TOKEN}}}"
f"{ETC_TOKEN}\n"
f"Tools: {serialized_tools}"
)
def build_chat_template_inputs(req: "ChatRequest"):
messages = normalize_messages_for_processor(req.messages)
tools, tool_choice_hint = select_tools_for_request(req.tools, req.tool_choice)
messages = prepend_system_message(messages, tool_choice_hint or "")
template_kwargs = {
"tokenize": True,
"add_generation_prompt": True,
"return_dict": True,
"return_tensors": "pt",
"enable_thinking": req.enable_thinking,
}
if tools:
template_kwargs["tools"] = tools
try:
inputs = processor.apply_chat_template(messages, **template_kwargs)
except TypeError as exc:
if not tools:
raise HTTPException(
status_code=400,
detail=f"apply_chat_template failed: {exc}",
) from exc
template_kwargs.pop("tools", None)
fallback_messages = prepend_system_message(messages, tool_prompt_fallback(tools))
try:
inputs = processor.apply_chat_template(
fallback_messages,
**template_kwargs,
)
except Exception as fallback_exc:
raise HTTPException(
status_code=400,
detail=f"apply_chat_template failed with tools: {fallback_exc}",
) from fallback_exc
except Exception as exc:
raise HTTPException(
status_code=400,
detail=f"apply_chat_template failed: {exc}",
) from exc
return inputs.to(model.device)
def arguments_to_json_string(arguments: Any) -> str:
if isinstance(arguments, str):
try:
json.loads(arguments)
return arguments
except json.JSONDecodeError:
arguments = parse_arguments_to_object(arguments)
try:
return json.dumps(arguments, ensure_ascii=False, separators=(",", ":"))
except TypeError:
return json.dumps(
{"value": str(arguments)},
ensure_ascii=False,
separators=(",", ":"),
)
def format_openai_tool_calls(
tool_calls: list[dict[str, Any]],
) -> list[dict[str, Any]]:
openai_tool_calls = []
for index, tool_call in enumerate(tool_calls):
function = tool_call.get("function", {})
name = function.get("name")
if not name:
continue
openai_tool_calls.append(
{
"id": tool_call.get("id") or f"call_{uuid.uuid4().hex[:24]}",
"type": "function",
"function": {
"name": str(name),
"arguments": arguments_to_json_string(
function.get("arguments", {})
),
},
}
)
return openai_tool_calls
def build_assistant_message(
thinking: str,
content: str,
tool_calls: list[dict[str, Any]],
*,
include_reasoning: bool = False,
finish_reason_override: str | None = None,
) -> tuple[dict[str, Any], str]:
text = select_visible_content(thinking, content)
openai_tool_calls = format_openai_tool_calls(tool_calls)
finish_reason = (
"tool_calls" if openai_tool_calls else finish_reason_override or "stop"
)
if openai_tool_calls:
message = {
"role": "assistant",
"content": text or None,
"tool_calls": openai_tool_calls,
}
else:
message = {
"role": "assistant",
"content": text or "模型未生成任何有效内容,请检查 num_inference_steps 或生成参数。",
}
if thinking and include_reasoning:
message["reasoning_content"] = thinking
return message, finish_reason
def length_finish_reason(
generated_len: int,
generation_kwargs: dict[str, Any],
) -> str | None:
max_new_tokens = generation_kwargs.get("max_new_tokens")
if isinstance(max_new_tokens, int) and generated_len >= max_new_tokens:
return "length"
return None
class ChatRequest(BaseModel):
model: str = SERVED_MODEL_NAME
messages: list[dict[str, Any]]
max_tokens: int | None = None
max_completion_tokens: int | None = DEFAULT_MAX_COMPLETION_TOKENS
num_inference_steps: int = DEFAULT_NUM_INFERENCE_STEPS
temperature: float | None = 0.4
stream: bool = False
enable_thinking: bool = False
tools: list[dict[str, Any]] | None = None
tool_choice: str | dict[str, Any] | None = None
parallel_tool_calls: bool | None = None
@app.post("/chat")
@app.post("/chat/completions")
@app.post("/v1/chat/completions")
async def chat(req: ChatRequest, request: Request):
inputs = build_chat_template_inputs(req)
generation_kwargs = build_generation_kwargs(inputs, req)
if req.stream:
async def stream_generator():
chat_id = f"chatcmpl-{int(time.time())}-{uuid.uuid4().hex[:8]}"
created = int(time.time())
yield "data: " + json.dumps({
"id": chat_id,
"object": "chat.completion.chunk",
"created": created,
"model": req.model,
"choices": [
{
"index": 0,
"delta": {"role": "assistant"},
"finish_reason": None,
}
],
}, ensure_ascii=False) + "\n\n"
with torch.no_grad():
outputs = model.generate(**generation_kwargs)
thinking, content, tool_calls, generated_len = extract_model_response(
outputs,
inputs,
)
full_text = select_visible_content(thinking, content)
openai_tool_calls = format_openai_tool_calls(tool_calls)
finish_reason = (
"tool_calls"
if openai_tool_calls
else length_finish_reason(generated_len, generation_kwargs) or "stop"
)
for i in range(0, len(full_text), 8):
chunk = full_text[i:i + 8]
yield "data: " + json.dumps({
"id": chat_id,
"object": "chat.completion.chunk",
"created": created,
"model": req.model,
"choices": [
{
"index": 0,
"delta": {"content": chunk},
"finish_reason": None,
}
],
}, ensure_ascii=False) + "\n\n"
await asyncio.sleep(0.02)
if openai_tool_calls:
yield "data: " + json.dumps({
"id": chat_id,
"object": "chat.completion.chunk",
"created": created,
"model": req.model,
"choices": [
{
"index": 0,
"delta": {
"tool_calls": [
{
"index": index,
**tool_call,
}
for index, tool_call in enumerate(
openai_tool_calls
)
]
},
"finish_reason": None,
}
],
}, ensure_ascii=False) + "\n\n"
await asyncio.sleep(0.02)
yield "data: " + json.dumps({
"id": chat_id,
"object": "chat.completion.chunk",
"created": created,
"model": req.model,
"choices": [
{
"index": 0,
"delta": {},
"finish_reason": finish_reason,
}
],
}, ensure_ascii=False) + "\n\n"
yield "data: [DONE]\n\n"
return StreamingResponse(stream_generator(), media_type="text/event-stream")
with torch.no_grad():
outputs = model.generate(**generation_kwargs)
prompt_len = inputs["input_ids"].shape[-1]
thinking, content, tool_calls, generated_len = extract_model_response(
outputs,
inputs,
)
message, finish_reason = build_assistant_message(
thinking,
content,
tool_calls,
include_reasoning=req.enable_thinking,
finish_reason_override=length_finish_reason(
generated_len,
generation_kwargs,
),
)
response = {
"id": f"chatcmpl-{int(time.time())}-{uuid.uuid4().hex[:8]}",
"object": "chat.completion",
"created": int(time.time()),
"model": req.model,
"choices": [
{
"index": 0,
"message": message,
"finish_reason": finish_reason,
}
],
"usage": {
"prompt_tokens": prompt_len,
"completion_tokens": generated_len,
"total_tokens": prompt_len + generated_len,
},
}
return response
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=9999)
需要更长的输出,则转头去改"generation_config.json",把"max_new_tokens"改为一个小于262144的值
为了兼容现在的OpenAI格式客户端,速度计算是不准的,例如New-API就会把速度记录为客户端流式渲染的速度,而不是服务端diffsion绘图的速度。比较蛋疼,但这个兼容起来工作量太大了,技术力不足(我太菜了


最新回复 (2)
-
斓曦未央丶 06-12 18:291楼感觉佬马上就可以转行干算法工程师走上人生巅峰了!
-
Cimix 楼主 06-13 00:102楼已经快被辞退了,天天被同事背后捅刀子,天天骂领导,时刻处于被边缘化
* 帖子来源Linux.do
附近帖子
- ↑哪里有强cpu+强gpu的算力租赁网站
- ↑谷歌云快到期了怎么办?
- ↑请问各位佬友如何提升英文水平?
- ↑一张图彻底明白Claude和ChatGPT订阅到底有多少额度,别听人忽悠了
- ↑用gpt-image-2去除图片噪点
- 📍 用Transformers部署DiffusionGemma,提供OpenAI格式接口
- ↓豆包下线思考模型,并新增「办公任务」模型(agent)
- ↓感谢any,fable5是真NB
- ↓新的抽奖帖出现了,不知道这次情况怎么样
- ↓【抽奖】两个TG号
- ↓ai自动化漏洞挖掘