用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:29
    1

    感觉佬马上就可以转行干算法工程师走上人生巅峰了!

  • Cimix 楼主 06-13 00:10
    2

    已经快被辞退了,天天被同事背后捅刀子,天天骂领导,时刻处于被边缘化

* 帖子来源Linux.do
返回