acdir-llada-math500 / tools /acdir_openai_server.py
NYCU-MLLab's picture
Upload folder using huggingface_hub
4a28d4d verified
Raw
History Blame Contribute Delete
30 kB
#!/usr/bin/env python
"""OpenAI-compatible ACDiR inference server for OpenCompass.
This server intentionally does not use LMDeploy's vanilla ``serve api_server``
target path. It exposes the small subset of OpenAI APIs needed by
OpenCompass while routing generation through the repository's ACDiR
critic-guided inference implementation.
"""
from __future__ import annotations
import argparse
import json
import os
import re
import sys
import threading
import time
import uuid
from pathlib import Path
from typing import Any
import torch
import uvicorn
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from transformers import AutoTokenizer
REPO_ROOT = Path(__file__).resolve().parents[1]
LMDEPLOY_ROOT = REPO_ROOT / "lmdeploy"
for path in (REPO_ROOT, LMDEPLOY_ROOT):
if path.exists() and str(path) not in sys.path:
sys.path.insert(0, str(path))
from metrics.phase2_critic_guided_math import ( # noqa: E402
_load_actor_adapter,
_load_critic_training_options,
_load_lmdeploy_actor_forward,
_merge_lora_adapter_into_lmdeploy_actor,
_normalize_llada_fast_mode,
_safe_torch_load,
_visible_set_build_kwargs_from_options,
generate_with_critic,
)
from networks.acdir import build_acdir_critic_from_actor # noqa: E402
from networks.modeling_llada import LLaDAModelLM # noqa: E402
_STATS_WRITE_LOCK = threading.Lock()
class ChatRequest(BaseModel):
model: str | None = None
messages: list[dict[str, Any]]
temperature: float | None = None
max_tokens: int | None = None
n: int | None = 1
def _message_text(message: dict[str, Any]) -> str:
content = message.get("content", "")
if isinstance(content, str):
return content
if isinstance(content, list):
chunks = []
for item in content:
if isinstance(item, dict):
if item.get("type") == "text":
chunks.append(str(item.get("text", "")))
elif "text" in item:
chunks.append(str(item.get("text", "")))
elif item is not None:
chunks.append(str(item))
return "\n".join(chunk for chunk in chunks if chunk)
return str(content)
def _prompt_from_messages(messages: list[dict[str, Any]]) -> str:
# OpenCompass encodes few-shot code tasks as alternating HUMAN/BOT chat
# messages. Dropping BOT turns turns the prompt into several unanswered
# user tasks, which makes MBPP generation copy the first demonstration.
return "\n".join(part for msg in messages if (part := _message_text(msg)))
def _sum_value(value: Any) -> float:
if value is None:
return 0.0
if isinstance(value, torch.Tensor):
if value.numel() == 0:
return 0.0
return float(value.detach().to(device="cpu", dtype=torch.float32).sum().item())
if isinstance(value, (list, tuple)):
return float(sum(_sum_value(item) for item in value))
try:
return float(value)
except (TypeError, ValueError):
return 0.0
def _total_at(mask_stats: dict[str, Any] | None, index: int) -> float:
if not isinstance(mask_stats, dict):
return 0.0
totals = mask_stats.get("totals")
if isinstance(totals, torch.Tensor) and totals.numel() > index:
return float(totals.detach().flatten()[index].to(device="cpu", dtype=torch.float32).item())
per_sample_totals = mask_stats.get("per_sample_totals")
if isinstance(per_sample_totals, torch.Tensor) and per_sample_totals.numel() > 0:
flat = per_sample_totals.detach().to(device="cpu", dtype=torch.float32)
if flat.ndim >= 2 and flat.shape[-1] > index:
return float(flat[:, index].sum().item())
return 0.0
def _summarize_mask_stats(
mask_stats: dict[str, Any] | None,
*,
remask_method: str,
reforward_after_remask: bool,
steps: int,
gen_length: int,
block_length: int,
block_steps: int | None,
llada_fast_mode: str,
count_logit_bias: str,
deterministic_joint_argmax: bool,
force_remask_window: int,
remask_candidate_disagree_only: bool,
remask_candidate_max_confidence: float,
) -> dict[str, Any]:
unique_remask = _total_at(mask_stats, 0)
denom_tokens = _total_at(mask_stats, 1)
remask = _total_at(mask_stats, 2)
same_step_refill = _total_at(mask_stats, 3)
same_step_same = _total_at(mask_stats, 4)
deferred_refill = _total_at(mask_stats, 5)
effective_change = _total_at(mask_stats, 6)
deferred_same = _total_at(mask_stats, 7)
deferred_pending = _total_at(mask_stats, 8)
refill = same_step_refill + deferred_refill
same_token = same_step_same + deferred_same
policy_decisions = _sum_value((mask_stats or {}).get("policy_decisions"))
if policy_decisions <= 0.0:
policy_decisions = _sum_value((mask_stats or {}).get("trajectory_action_count"))
count0 = _sum_value((mask_stats or {}).get("count0_decisions"))
count1 = _sum_value((mask_stats or {}).get("count1_decisions"))
count2 = _sum_value((mask_stats or {}).get("count2_decisions"))
selected_score_sum = _sum_value((mask_stats or {}).get("selected_score_sum"))
selected_score_count = _sum_value((mask_stats or {}).get("selected_score_count"))
action_score_sum = _sum_value((mask_stats or {}).get("trajectory_action_score_sum"))
action_score_count = _sum_value((mask_stats or {}).get("trajectory_action_count"))
candidate_count_sum = _sum_value((mask_stats or {}).get("candidate_count_sum"))
return {
"remask_method": str(remask_method),
"reforward_after_remask": bool(reforward_after_remask),
"steps": int(steps),
"gen_length": int(gen_length),
"block_length": int(block_length),
"block_steps": None if block_steps is None else int(block_steps),
"llada_fast_mode": str(llada_fast_mode),
"count_logit_bias": str(count_logit_bias or ""),
"deterministic_joint_argmax": bool(deterministic_joint_argmax),
"force_remask_window": int(force_remask_window or 0),
"remask_candidate_disagree_only": bool(remask_candidate_disagree_only),
"remask_candidate_max_confidence": float(remask_candidate_max_confidence or 0.0),
"unique_remask_count": unique_remask,
"remask_count": remask,
"remask_ratio": unique_remask / max(denom_tokens, 1.0),
"policy_decisions": policy_decisions,
"count0_decisions": count0,
"count1_decisions": count1,
"count2_decisions": count2,
"candidate_count_mean": candidate_count_sum / max(policy_decisions, 1.0),
"same_step_refill_count": same_step_refill,
"deferred_refill_count": deferred_refill,
"deferred_pending_count": deferred_pending,
"same_token_refill_count": same_token,
"effective_change_count": effective_change,
"same_token_refill_rate": same_token / max(refill, 1.0),
"effective_change_rate": effective_change / max(refill, 1.0),
"selected_score_mean": selected_score_sum / max(selected_score_count, 1.0),
"action_score_mean": action_score_sum / max(action_score_count, 1.0),
}
def _format_stats_line(stats: dict[str, Any]) -> str:
return (
f"method={stats['remask_method']} reforward={stats['reforward_after_remask']} "
f"fast={stats.get('llada_fast_mode', 'full_window')} "
f"joint={stats.get('deterministic_joint_argmax', False)} win={stats.get('force_remask_window', 0)} "
f"gate={stats.get('remask_candidate_disagree_only', False)}/{stats.get('remask_candidate_max_confidence', 0.0):.2f} "
f"remask={stats['remask_count']:.1f} unique={stats['unique_remask_count']:.1f} "
f"rho={stats['remask_ratio'] * 100:.2f}% policy={stats['policy_decisions']:.1f} "
f"k0/1/2={stats['count0_decisions']:.0f}/{stats['count1_decisions']:.0f}/{stats['count2_decisions']:.0f} "
f"cand={stats['candidate_count_mean']:.1f} refillS/D={stats['same_step_refill_count']:.1f}/{stats['deferred_refill_count']:.1f} "
f"sameTok={stats['same_token_refill_rate'] * 100:.1f}% effChg={stats['effective_change_rate'] * 100:.1f}% "
f"pending={stats['deferred_pending_count']:.1f} score={stats['selected_score_mean']:.3f}"
)
def _append_stats_jsonl(stats: dict[str, Any]) -> None:
stats_path = os.environ.get("ACDIR_SERVER_STATS_JSONL", "").strip()
if not stats_path:
return
path = Path(stats_path)
path.parent.mkdir(parents=True, exist_ok=True)
line = json.dumps(stats, ensure_ascii=False, sort_keys=True) + "\n"
with _STATS_WRITE_LOCK:
with path.open("a", encoding="utf-8") as handle:
handle.write(line)
def _env_bool(name: str, default: bool = False) -> bool:
value = os.environ.get(name)
if value is None or value == "":
return default
return value.strip().lower() in {"1", "true", "yes", "y", "on"}
def _normalize_repetition_key(text: str) -> str:
text = re.sub(r"\s+", " ", str(text or "").lower()).strip()
text = re.sub(r"[^a-z0-9]+", " ", text)
return re.sub(r"\s+", " ", text).strip()
def _truncate_summary_repetition(text: str, *, max_sentences: int = 6, min_sentence_words: int = 6) -> str:
text = re.sub(r"\s+", " ", str(text or "")).strip()
if not text:
return text
text = re.sub(r"\s+([,.;:!?])", r"\1", text)
text = re.sub(r"([([{])\s+", r"\1", text)
filler = re.search(r"(?:\{?\s*(?:vocalsound|laugh|noise|silence)\s*\}?\s*){3,}", text, flags=re.IGNORECASE)
if filler is not None:
text = text[: filler.start()].strip()
text = re.sub(r"\b[A-Z][A-Za-z ]{0,40}:\s*$", "", text).strip()
if not text:
return ""
raw_tokens = text.split()
norm_pairs = [
(raw_idx, norm)
for raw_idx, token in enumerate(raw_tokens)
for norm in [_normalize_repetition_key(token)]
if norm
]
# Diffusion decoding can fall into short phrase cycles such as
# "play, pause, stop, play, pause, stop". Sentence-level filtering does
# not catch these because the loop often has no final punctuation.
norms = [norm for _, norm in norm_pairs]
for n in range(1, min(13, max(1, len(norms) // 2 + 1))):
min_repeats = 3 if n <= 6 else 2
max_start = len(norms) - (n * min_repeats)
for start in range(max_start + 1):
span = norms[start : start + n]
if not span:
continue
if all(norms[start + rep * n : start + (rep + 1) * n] == span for rep in range(1, min_repeats)):
raw_cut = norm_pairs[start + n][0]
text = " ".join(raw_tokens[:raw_cut]).strip()
raw_tokens = text.split()
break
else:
continue
break
if not text:
return ""
repeat_run = 1
prev_norm = ""
for idx, token in enumerate(raw_tokens):
norm = _normalize_repetition_key(token)
if norm and norm == prev_norm:
repeat_run += 1
else:
repeat_run = 1
prev_norm = norm
if repeat_run >= 6:
text = " ".join(raw_tokens[: idx - repeat_run + 1]).strip()
break
if not text:
return ""
pieces = [part.strip() for part in re.split(r"(?<=[.!?])\s+", text) if part.strip()]
if not pieces:
pieces = [text]
kept: list[str] = []
seen: set[str] = set()
prev_key = ""
for piece in pieces:
key = _normalize_repetition_key(piece)
words = key.split()
if len(words) >= int(min_sentence_words):
if key in seen and kept:
break
if prev_key:
prev = set(prev_key.split())
cur = set(words)
overlap = len(prev & cur) / max(1, min(len(prev), len(cur)))
if overlap >= 0.85 and kept:
break
kept.append(piece)
if len(words) >= int(min_sentence_words):
seen.add(key)
prev_key = key
if len(kept) >= int(max_sentences):
break
out = " ".join(kept).strip() or text
tokens = out.split()
seen_ngrams: dict[tuple[str, ...], int] = {}
for n in (4, 6, 8):
seen_ngrams.clear()
ngram_counts: dict[tuple[str, ...], int] = {}
for idx in range(0, max(0, len(tokens) - n + 1)):
gram = tuple(_normalize_repetition_key(" ".join(tokens[idx : idx + n])).split())
if len(gram) < n:
continue
first = seen_ngrams.setdefault(gram, idx)
ngram_counts[gram] = ngram_counts.get(gram, 0) + 1
threshold = 3 if n <= 6 else 2
if ngram_counts[gram] >= threshold and idx >= first + n:
out = " ".join(tokens[:idx]).strip()
tokens = out.split()
break
return out.strip()
class AcdirEngine:
def __init__(self, args: argparse.Namespace):
args.llada_fast_mode = _normalize_llada_fast_mode(args.llada_fast_mode)
self.args = args
self.actor_only = str(args.remask_method or "").strip().lower() in {
"actor_only",
"actor-only",
"none",
"baseline",
"noop",
"no_op",
"",
}
critic_name = Path(args.critic_ckpt_path).name if args.critic_ckpt_path else "actor-only"
self.model_id = args.model_id or f"acdir:{critic_name}"
self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
if self.device == "cpu":
raise RuntimeError("ACDiR OpenAI server requires a visible CUDA device.")
torch.set_grad_enabled(False)
torch.manual_seed(int(args.seed))
torch.cuda.set_device(0)
self.tokenizer = AutoTokenizer.from_pretrained(args.ckpt_path)
self.tokenizer.pad_token_id = self.tokenizer.eos_token_id
self.tokenizer.padding_side = "left"
self.actor_forward, self.actor_backend = self._load_actor()
self.critic = None if self.actor_only else self._load_critic()
requested_time_dim = int(args.time_embed_dim or 0)
critic_time_dim = int(getattr(self.critic, "time_embed_dim", 0) or 0)
self.time_embed_dim = requested_time_dim if requested_time_dim > 0 else critic_time_dim
if self.time_embed_dim <= 0:
self.time_embed_dim = int(
getattr(self.actor_forward.config, "hidden_size", None)
or getattr(self.actor_forward.config, "d_model", 0)
or 0
)
if critic_time_dim > 0 and int(self.time_embed_dim) != critic_time_dim:
raise ValueError(
f"time_embed_dim={self.time_embed_dim} does not match critic checkpoint "
f"time_embed_dim={critic_time_dim}; use --time_embed_dim 0 for auto."
)
self.lock = threading.Lock()
token_mask_id = self.tokenizer.convert_tokens_to_ids("<|mdm_mask|>")
self.mask_id = token_mask_id if token_mask_id is not None and token_mask_id >= 0 else int(args.mask_id)
self.eos_id = self.tokenizer.eos_token_id if int(args.eos_id) < 0 else int(args.eos_id)
print(
"[acdir-server] ready "
f"model_id={self.model_id} actor_backend={self.actor_backend} "
f"critic={args.critic_ckpt_path or '<disabled>'} device={self.device} "
f"steps={args.steps} gen_length={args.gen_length} block={args.block_length}/{args.block_steps} "
f"remask_method={args.remask_method} reforward={args.reforward_after_remask} "
f"fast_mode={args.llada_fast_mode} "
f"count_logit_bias={args.count_logit_bias or '<none>'} "
f"joint_argmax={args.deterministic_joint_argmax} force_window={args.force_remask_window} "
f"candidate_gate={args.remask_candidate_disagree_only}/{args.remask_candidate_max_confidence} "
f"respect_request_max_tokens={args.respect_request_max_tokens} "
f"time_embed_dim={self.time_embed_dim}",
flush=True,
)
def _load_actor(self):
backend = str(self.args.actor_forward_backend).lower()
if backend in {"lmdeploy", "auto"}:
try:
actor_forward = _load_lmdeploy_actor_forward(
self.args.ckpt_path,
device=self.device,
dtype=self.args.actor_forward_dtype,
)
if self.args.adapter_path:
merged_ckpt, merged_count = _merge_lora_adapter_into_lmdeploy_actor(
actor_forward,
self.args.adapter_path,
)
print(f"[acdir-server] merged actor LoRA {merged_count} tensors from {merged_ckpt}", flush=True)
return actor_forward, "lmdeploy"
except Exception:
if backend == "lmdeploy":
raise
print("[acdir-server] LMDeploy actor unavailable; falling back to HF actor.", flush=True)
torch.cuda.empty_cache()
actor = LLaDAModelLM.from_pretrained(
pretrained_model_name_or_path=self.args.ckpt_path,
torch_dtype=torch.bfloat16,
)
if self.args.adapter_path:
actor = _load_actor_adapter(actor, self.args.adapter_path)
actor.eval().requires_grad_(False).to(self.device)
actor.config.use_cache = False
if hasattr(actor, "generation_config") and actor.generation_config is not None:
actor.generation_config.use_cache = False
return actor, "hf"
def _load_critic(self):
if not self.args.critic_ckpt_path:
raise ValueError("--critic_ckpt_path is required unless --remask_method actor_only is used.")
options = _load_critic_training_options(self.args.critic_ckpt_path)
critic_kwargs = _visible_set_build_kwargs_from_options(options)
print(f"[acdir-server] critic build kwargs: {critic_kwargs}", flush=True)
critic = build_acdir_critic_from_actor(self.actor_forward, **critic_kwargs)
state = _safe_torch_load(self.args.critic_ckpt_path, map_location="cpu")
critic.load_state_dict(state, strict=True)
critic.eval().to(self.device)
return critic
def generate_one(
self,
prompt: str,
request_temperature: float | None = None,
request_max_tokens: int | None = None,
) -> tuple[str, dict[str, Any]]:
if self.args.respect_request_temperature and request_temperature is not None:
temperature = float(request_temperature)
no_sample = temperature <= 0.0
else:
temperature = float(self.args.temperature)
no_sample = bool(self.args.no_sample)
gen_length = int(self.args.gen_length)
block_length = int(self.args.block_length)
steps = int(self.args.steps)
block_steps = None if int(self.args.block_steps or 0) <= 0 else int(self.args.block_steps)
if self.args.respect_request_max_tokens and request_max_tokens is not None and int(request_max_tokens) > 0:
gen_length = int(request_max_tokens)
if gen_length % block_length != 0:
block_length = gen_length
if steps == int(self.args.gen_length):
steps = gen_length
if block_steps is not None and block_steps > block_length:
block_steps = block_length
generation_started = time.time()
with self.lock:
responses, _stats = generate_with_critic(
self.actor_forward,
self.critic,
self.tokenizer,
[prompt],
steps=steps,
gen_length=gen_length,
block_length=block_length,
block_steps=block_steps,
device=self.device,
no_sample=no_sample,
temperature=temperature,
cfg_scale=float(self.args.cfg_scale),
unmask_policy="confidence",
mask_id=self.mask_id,
eos_id=self.eos_id,
time_embed_dim=int(self.time_embed_dim),
use_chat_template=bool(self.args.use_chat_template),
prompt_style=str(self.args.prompt_style),
lmdeploy_cuda_graph=False,
llada_fast_mode=str(self.args.llada_fast_mode),
sdar_confidence_threshold=float(self.args.sdar_confidence_threshold),
remask_method=str(self.args.remask_method),
max_total_remask_per_sample=int(self.args.max_total_remask_per_sample),
reforward_after_remask=bool(self.args.reforward_after_remask),
count_logit_bias=str(self.args.count_logit_bias or ""),
remask_candidate_disagree_only=bool(self.args.remask_candidate_disagree_only),
remask_candidate_max_confidence=float(self.args.remask_candidate_max_confidence or 0.0),
deterministic_joint_argmax=bool(self.args.deterministic_joint_argmax),
force_remask_window=int(self.args.force_remask_window or 0),
logits_eos_inf=bool(self.args.logits_eos_inf),
confidence_eos_eot_inf=bool(self.args.confidence_eos_eot_inf),
)
text = responses[0] if responses else ""
raw_text = text
if bool(self.args.summary_truncate_repetition):
text = _truncate_summary_repetition(
text,
max_sentences=int(self.args.summary_truncate_max_sentences),
)
stats = _summarize_mask_stats(
_stats,
remask_method=str(self.args.remask_method),
reforward_after_remask=bool(self.args.reforward_after_remask),
steps=steps,
gen_length=gen_length,
block_length=block_length,
block_steps=block_steps,
llada_fast_mode=str(self.args.llada_fast_mode),
count_logit_bias=str(self.args.count_logit_bias or ""),
deterministic_joint_argmax=bool(self.args.deterministic_joint_argmax),
force_remask_window=int(self.args.force_remask_window or 0),
remask_candidate_disagree_only=bool(self.args.remask_candidate_disagree_only),
remask_candidate_max_confidence=float(self.args.remask_candidate_max_confidence or 0.0),
)
stats["generation_elapsed_s"] = time.time() - generation_started
if bool(self.args.summary_truncate_repetition):
stats["postprocess_summary_truncate_repetition"] = True
stats["raw_output_chars"] = int(len(raw_text))
stats["postprocess_removed_chars"] = int(max(0, len(raw_text) - len(text)))
return text, stats
def build_app(engine: AcdirEngine) -> FastAPI:
app = FastAPI()
@app.get("/v1/models")
def models():
return {
"object": "list",
"data": [
{
"id": engine.model_id,
"object": "model",
"created": 0,
"owned_by": "acdir",
}
],
}
@app.post("/v1/chat/completions")
def chat_completions(req: ChatRequest):
prompt = _prompt_from_messages(req.messages)
if not prompt:
raise HTTPException(status_code=400, detail="empty prompt")
request_id = uuid.uuid4().hex
preview_chars = int(os.environ.get("ACDIR_SERVER_PROMPT_PREVIEW_CHARS", "0") or 0)
if preview_chars > 0:
preview = prompt[:preview_chars].replace("\n", "\\n")
print(f"[acdir-server] prompt preview: {preview}", flush=True)
n = max(1, int(req.n or 1))
choices = []
started = time.time()
request_log = _env_bool("ACDIR_SERVER_REQUEST_LOG", True)
if request_log:
print(
f"[acdir-server] request start id={request_id} n={n} prompt_chars={len(prompt)} "
f"temperature={req.temperature} max_tokens={req.max_tokens}",
flush=True,
)
for idx in range(n):
text, stats = engine.generate_one(prompt, req.temperature, req.max_tokens)
stats.update(
{
"request_id": request_id,
"item_index": int(idx),
"n": int(n),
"prompt_chars": int(len(prompt)),
"output_chars": int(len(text)),
"elapsed_s": time.time() - started,
"timestamp": time.time(),
}
)
_append_stats_jsonl(stats)
choices.append(
{
"index": idx,
"message": {"role": "assistant", "content": text},
"finish_reason": "stop",
}
)
if request_log:
print(
f"[acdir-server] request item done id={request_id} index={idx} output_chars={len(text)} "
f"elapsed={time.time() - started:.1f}s {_format_stats_line(stats)}",
flush=True,
)
now = int(time.time())
if request_log:
print(
f"[acdir-server] request done id={request_id} n={n} elapsed={time.time() - started:.1f}s",
flush=True,
)
return {
"id": f"chatcmpl-acdir-{request_id}",
"object": "chat.completion",
"created": now,
"model": engine.model_id,
"choices": choices,
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0,
},
}
return app
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Serve ACDiR as an OpenAI-compatible target model.")
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=23333)
parser.add_argument("--model_id", default="")
parser.add_argument("--ckpt_path", required=True)
parser.add_argument("--critic_ckpt_path", default="")
parser.add_argument("--adapter_path", default="")
parser.add_argument("--steps", type=int, default=256)
parser.add_argument("--gen_length", type=int, default=512)
parser.add_argument("--block_length", type=int, default=32)
parser.add_argument("--block_steps", type=int, default=16)
parser.add_argument("--temperature", type=float, default=0.0)
parser.add_argument("--no_sample", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--respect_request_temperature", action="store_true")
parser.add_argument("--respect_request_max_tokens", action="store_true")
parser.add_argument("--cfg_scale", type=float, default=0.0)
parser.add_argument("--mask_id", type=int, default=126336)
parser.add_argument("--eos_id", type=int, default=126081)
parser.add_argument("--remask_method", default="count_set")
parser.add_argument("--logits_eos_inf", action=argparse.BooleanOptionalAction, default=False)
parser.add_argument("--confidence_eos_eot_inf", action=argparse.BooleanOptionalAction, default=False)
parser.add_argument("--max_total_remask_per_sample", type=int, default=0)
parser.add_argument("--reforward_after_remask", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--count_logit_bias", default="", help="Comma-separated inference bias for count logits, e.g. 0,0.6,0.9.")
parser.add_argument(
"--remask_candidate_disagree_only",
action=argparse.BooleanOptionalAction,
default=False,
help="Only allow visible tokens whose current actor argmax disagrees with the committed token to be remask candidates.",
)
parser.add_argument(
"--remask_candidate_max_confidence",
type=float,
default=0.0,
help="Also allow remask candidates whose committed-token confidence is below this threshold; 0 disables.",
)
parser.add_argument("--deterministic_joint_argmax", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--force_remask_window", type=int, default=1, help="Force count-set window id; <=0 leaves the critic window head active.")
parser.add_argument("--time_embed_dim", type=int, default=0)
parser.add_argument("--sdar_confidence_threshold", type=float, default=0.85)
parser.add_argument(
"--llada_fast_mode",
default=os.environ.get("LLADA_LMDEPLOY_FAST_MODE", "full_window"),
choices=["full_window", "full-window", "full", "exact", "block_cache", "block-cache", "prefix_cache", "prefix-cache", "lmdeploy_cache", "truncate_future", "truncate-future"],
help="LLaDA inference fast path. block_cache reuses prompt/completed-block KV cache and is approximate.",
)
parser.add_argument("--actor_forward_backend", choices=["hf", "lmdeploy", "auto"], default="hf")
parser.add_argument("--actor_forward_dtype", default="bfloat16")
parser.add_argument("--use_chat_template", action=argparse.BooleanOptionalAction, default=True)
parser.add_argument("--prompt_style", default="default")
parser.add_argument("--seed", type=int, default=113)
parser.add_argument("--summary_truncate_repetition", action=argparse.BooleanOptionalAction, default=False)
parser.add_argument("--summary_truncate_max_sentences", type=int, default=6)
return parser.parse_args()
def main() -> None:
args = parse_args()
engine = AcdirEngine(args)
app = build_app(engine)
uvicorn.run(app, host=args.host, port=args.port, log_level=os.environ.get("ACDIR_SERVER_LOG_LEVEL", "info"))
if __name__ == "__main__":
main()