kingjones777's picture
Add files using upload-large-folder tool
18c1466 verified
Raw
History Blame Contribute Delete
17.3 kB
#!/usr/bin/env python3
"""Prompt enhancement (PE) for Ming-Image text-to-image via a Ling-3.0-flash-VL seat.
Per the README, PE is a pre-processing step *outside* ``infer.py``: an
instruction-following VLM rewrites a short caption into the structured
Figma-style JSON prompt that the text-to-image pipeline consumes, and the
result is passed to ``infer.py --prompt`` as raw text or via a file.
This module drives any OpenAI-compatible ``/chat/completions`` endpoint using
only the standard library (``urllib``): by default the local llama-server seat
serving Ling-3.0-flash-VL, optionally the LiteLLM lab gateway (Bearer auth via
``--api-key`` or the ``LITELLM_API_KEY`` environment variable). The rewriter
system prompt is read verbatim from ``assets/t2i_rewriter_system_prompt.txt``.
The reply is parsed robustly (```json fences and surrounding prose are
tolerated), then validated against the schema the system prompt demands. On a
parse or validation failure the request is retried exactly once with the
errors appended to the user turn; if that still fails, PromptEnhancementError
is raised with the errors. Invalid JSON is never passed through silently.
CLI:
python pe_ling.py "a caption" --out prompt.json \
[--base-url http://127.0.0.1:8090/v1] \
[--model ling-3.0-flash-vl-mtp-halo-STRIX_LEAN]
"""
from __future__ import annotations
import argparse
import json
import os
import re
import sys
import time
import urllib.error
import urllib.request
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
CODE_DIRECTORY = Path(__file__).resolve().parent
SYSTEM_PROMPT_PATH = CODE_DIRECTORY / "assets" / "t2i_rewriter_system_prompt.txt"
# The Ling-3.0-flash-VL seat already served on the target box (llama-server,
# OpenAI-compatible, thinking disabled); both endpoints speak the same
# /chat/completions protocol.
DEFAULT_BASE_URL = "http://127.0.0.1:8090/v1"
DEFAULT_MODEL = "ling-3.0-flash-vl-mtp-halo-STRIX_LEAN"
API_KEY_ENV = "LITELLM_API_KEY"
# Low temperature: the rewrite is a deterministic schema transformation, not
# creative sampling.
DEFAULT_TEMPERATURE = 0.2
# The upstream example rewrite (assets/t2i_four_seasons_cabin_prompt.json) is
# ~5 KB (~2k tokens); dense multi-layer infographic rewrites run several times
# longer, so leave generous headroom for a complete JSON object.
DEFAULT_MAX_TOKENS = 16384
# A multi-thousand-token completion on the local seat can take minutes.
DEFAULT_TIMEOUT_SECONDS = 600.0
REPAIR_INSTRUCTION = "Return only the corrected JSON object: no prose, no code fences."
CANVAS_SETTINGS_KEYS = ("aspect_ratio", "ambient_lighting", "image_style")
LAYER_KEYS = ("description", "coordinates", "hierarchy_and_relation", "color_specs")
COORDINATE_FIELDS = ("cx", "cy", "w", "h")
# `coordinates` must be ONE string of the form
# "cx: 0.500, cy: 0.500, w: 1.000, h: 1.000". The upstream example also uses
# bare integers ("h: 1"), so accept any decimal spelling and enforce the
# [0, 1] range on the parsed value. Whitespace around ':' and ',' is
# tolerated; the key order is fixed.
_COORDINATE_NUMBER = r"[-+]?(?:\d+(?:\.\d*)?|\.\d+)"
COORDINATES_RE = re.compile(
rf"^\s*cx:\s*(?P<cx>{_COORDINATE_NUMBER})\s*,"
rf"\s*cy:\s*(?P<cy>{_COORDINATE_NUMBER})\s*,"
rf"\s*w:\s*(?P<w>{_COORDINATE_NUMBER})\s*,"
rf"\s*h:\s*(?P<h>{_COORDINATE_NUMBER})\s*$"
)
# Hex colors: #RGB, #RGBA, #RRGGBB, #RRGGBBAA (the upstream example uses
# #RRGGBB; the alpha forms keep RGBA-design outputs from failing validation).
HEX_COLOR_RE = re.compile(
r"^#(?:[0-9a-fA-F]{3}|[0-9a-fA-F]{4}|[0-9a-fA-F]{6}|[0-9a-fA-F]{8})$"
)
class PromptEnhancementError(RuntimeError):
"""PE failed: transport/protocol error, or schema failure after the retry."""
def __init__(
self,
message: str,
errors: Optional[List[str]] = None,
reply: Optional[str] = None,
):
super().__init__(message)
self.errors = list(errors or [])
self.reply = reply
def load_system_prompt(path: Path = SYSTEM_PROMPT_PATH) -> str:
"""Return the released rewriter system prompt, verbatim."""
return path.read_text(encoding="utf-8")
def extract_json_object(text: str) -> Dict[str, Any]:
"""Return the first complete top-level JSON object found in ``text``.
Models sometimes wrap JSON in ```json fences or add prose around it.
Scanning every ``{`` position with ``JSONDecoder.raw_decode`` (which
decodes a document at an offset and ignores trailing data) recovers the
object in all of those shapes. Raises ValueError when no complete JSON
object is present, e.g. a reply truncated mid-object.
"""
decoder = json.JSONDecoder()
position = text.find("{")
while position != -1:
try:
document, _ = decoder.raw_decode(text, position)
except ValueError:
position = text.find("{", position + 1)
continue
return document
snippet = text.strip()
if len(snippet) > 300:
snippet = snippet[:300] + "..."
raise ValueError(
f"reply contains no complete top-level JSON object "
f"({len(text)} characters); starts with: {snippet!r}"
)
def _check_exact_keys(
mapping: Dict[str, Any], expected: Tuple[str, ...], path: str, errors: List[str]
) -> None:
missing = [key for key in expected if key not in mapping]
unexpected = [key for key in mapping if key not in expected]
if missing:
errors.append(f"{path}: missing required key(s): {', '.join(missing)}")
if unexpected:
errors.append(
f"{path}: unexpected key(s): {', '.join(unexpected)} "
f"(exactly {', '.join(expected)} are required)"
)
def _check_non_empty_string(value: Any, path: str, errors: List[str]) -> None:
if not isinstance(value, str):
errors.append(f"{path}: expected a string, got {type(value).__name__}")
elif not value.strip():
errors.append(f"{path}: string is empty")
def _check_coordinates(value: Any, path: str, errors: List[str]) -> None:
if not isinstance(value, str):
errors.append(
f"{path}: must be ONE string of the form "
f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000', got {type(value).__name__}"
)
return
match = COORDINATES_RE.match(value)
if match is None:
errors.append(
f"{path}: {value!r} is not of the form "
f"'cx: 0.500, cy: 0.500, w: 1.000, h: 1.000'"
)
return
for field in COORDINATE_FIELDS:
number = float(match.group(field))
if not 0.0 <= number <= 1.0:
errors.append(f"{path}: {field}={match.group(field)} is outside [0, 1]")
def _check_color_specs(value: Any, path: str, errors: List[str]) -> None:
if not isinstance(value, list):
errors.append(
f"{path}: expected a list of hex colors, got {type(value).__name__}"
)
return
for index, color in enumerate(value):
if not isinstance(color, str) or HEX_COLOR_RE.match(color) is None:
errors.append(
f"{path}[{index}]: {color!r} is not a hex color "
f"(expected #RGB, #RGBA, #RRGGBB, or #RRGGBBAA)"
)
def validate_enhanced_prompt(document: Any) -> List[str]:
"""Return schema errors for a rewritten prompt; an empty list means valid.
Schema demanded by assets/t2i_rewriter_system_prompt.txt: exactly two
top-level keys ``canvas_settings`` (exactly ``aspect_ratio``,
``ambient_lighting``, ``image_style``) and ``layers`` (each layer exactly
``description``, ``coordinates``, ``hierarchy_and_relation``,
``color_specs``); ``coordinates`` is a string "cx: 0.500, cy: 0.500,
w: 1.000, h: 1.000" with values in [0, 1]; ``color_specs`` is a list of
hex colors. ``layers`` must hold at least one visible layer -- an empty
list means the rewrite failed even though it is type-correct.
"""
if not isinstance(document, dict):
return [f"top level: expected a JSON object, got {type(document).__name__}"]
errors: List[str] = []
_check_exact_keys(document, ("canvas_settings", "layers"), "top level", errors)
if "canvas_settings" in document:
canvas = document["canvas_settings"]
if not isinstance(canvas, dict):
errors.append(
f"canvas_settings: expected a JSON object, got {type(canvas).__name__}"
)
else:
_check_exact_keys(canvas, CANVAS_SETTINGS_KEYS, "canvas_settings", errors)
for key in CANVAS_SETTINGS_KEYS:
if key in canvas:
_check_non_empty_string(
canvas[key], f"canvas_settings.{key}", errors
)
if "layers" in document:
layers = document["layers"]
if not isinstance(layers, list):
errors.append(f"layers: expected a list, got {type(layers).__name__}")
elif not layers:
errors.append("layers: expected at least one visible layer")
else:
for index, layer in enumerate(layers):
path = f"layers[{index}]"
if not isinstance(layer, dict):
errors.append(
f"{path}: expected a JSON object, got {type(layer).__name__}"
)
continue
_check_exact_keys(layer, LAYER_KEYS, path, errors)
for key in ("description", "hierarchy_and_relation"):
if key in layer:
_check_non_empty_string(layer[key], f"{path}.{key}", errors)
if "coordinates" in layer:
_check_coordinates(
layer["coordinates"], f"{path}.coordinates", errors
)
if "color_specs" in layer:
_check_color_specs(layer["color_specs"], f"{path}.color_specs", errors)
return errors
def _chat_completion(
base_url: str,
model: str,
messages: List[Dict[str, str]],
*,
temperature: float,
max_tokens: int,
api_key: Optional[str],
timeout: float,
) -> Tuple[str, Optional[str]]:
"""POST one chat completion; return (content, finish_reason)."""
url = base_url.rstrip("/") + "/chat/completions"
payload = json.dumps(
{
"model": model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
"stream": False,
}
).encode("utf-8")
headers = {"Content-Type": "application/json"}
if api_key:
headers["Authorization"] = f"Bearer {api_key}"
request = urllib.request.Request(url, data=payload, headers=headers, method="POST")
try:
with urllib.request.urlopen(request, timeout=timeout) as response:
body = response.read().decode("utf-8", errors="replace")
except urllib.error.HTTPError as error:
detail = error.read().decode("utf-8", errors="replace")
raise PromptEnhancementError(
f"HTTP {error.code} from {url}: {detail[:2000]}"
) from error
except urllib.error.URLError as error:
raise PromptEnhancementError(f"cannot reach {url}: {error.reason}") from error
except OSError as error: # includes socket timeouts during the read
raise PromptEnhancementError(f"request to {url} failed: {error}") from error
try:
envelope = json.loads(body)
choice = envelope["choices"][0]
content = choice["message"]["content"]
except (json.JSONDecodeError, KeyError, IndexError, TypeError) as error:
raise PromptEnhancementError(
f"malformed chat completion response from {url}: {body[:500]}"
) from error
finish_reason = choice.get("finish_reason")
if not isinstance(content, str) or not content.strip():
raise PromptEnhancementError(
f"empty completion content from {url} (finish_reason={finish_reason!r})"
)
return content, finish_reason
def enhance(
caption: str,
base_url: str,
model: str,
api_key: Optional[str] = None,
timeout: float = DEFAULT_TIMEOUT_SECONDS,
temperature: float = DEFAULT_TEMPERATURE,
max_tokens: int = DEFAULT_MAX_TOKENS,
) -> Dict[str, Any]:
"""Return the validated structured rewrite of ``caption``.
Sends the verbatim rewriter system prompt plus the caption to
``{base_url}/chat/completions``. On a parse or schema failure, retries
exactly once with the validation errors appended to the user turn; if
that also fails, raises PromptEnhancementError carrying the errors.
"""
system_prompt = load_system_prompt()
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": caption},
]
request_kwargs = {
"temperature": temperature,
"max_tokens": max_tokens,
"api_key": api_key,
"timeout": timeout,
}
errors: List[str] = []
content = ""
for attempt in (1, 2):
content, finish_reason = _chat_completion(
base_url, model, messages, **request_kwargs
)
document: Optional[Dict[str, Any]] = None
try:
document = extract_json_object(content)
except ValueError as error:
errors = [str(error)]
if document is not None:
errors = validate_enhanced_prompt(document)
if not errors:
assert document is not None # errors empty implies extraction succeeded
return document
if finish_reason == "length":
errors.append(
"the reply was cut off (finish_reason='length'): the complete "
f"JSON object must fit within max_tokens={max_tokens}"
)
print(f"pe_ling: attempt {attempt}/2 failed validation:", file=sys.stderr)
for error in errors:
print(f"pe_ling: - {error}", file=sys.stderr)
if attempt == 1:
retry_content = (
f"{caption}\n\n"
"Your previous reply failed schema validation:\n"
+ "".join(f"- {error}\n" for error in errors)
+ "\n"
+ REPAIR_INSTRUCTION
)
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": retry_content},
]
raise PromptEnhancementError(
"prompt enhancement failed schema validation after 2 attempts:\n"
+ "".join(f" - {error}\n" for error in errors).rstrip(),
errors=errors,
reply=content,
)
def main() -> None:
parser = argparse.ArgumentParser(
description=(
"Enhance a Ming-Image text-to-image caption into the validated "
"structured JSON prompt via an OpenAI-compatible Ling-3.0-flash-VL "
"endpoint."
)
)
parser.add_argument("caption", help="free-form design caption to enhance")
parser.add_argument(
"--out",
type=Path,
help="write the validated JSON here (default: stdout, summary on stderr)",
)
parser.add_argument(
"--base-url",
default=DEFAULT_BASE_URL,
help=f"OpenAI-compatible base URL (default: {DEFAULT_BASE_URL})",
)
parser.add_argument(
"--model",
default=DEFAULT_MODEL,
help=f"chat model id served at the endpoint (default: {DEFAULT_MODEL})",
)
parser.add_argument(
"--api-key",
default=os.environ.get(API_KEY_ENV),
help=f"Bearer token for gated endpoints; defaults to ${API_KEY_ENV} when set",
)
parser.add_argument(
"--timeout",
type=float,
default=DEFAULT_TIMEOUT_SECONDS,
help=f"per-request timeout in seconds (default: {DEFAULT_TIMEOUT_SECONDS})",
)
parser.add_argument(
"--temperature",
type=float,
default=DEFAULT_TEMPERATURE,
help=f"sampling temperature (default: {DEFAULT_TEMPERATURE})",
)
parser.add_argument(
"--max-tokens",
type=int,
default=DEFAULT_MAX_TOKENS,
help=f"completion token budget (default: {DEFAULT_MAX_TOKENS})",
)
args = parser.parse_args()
started = time.perf_counter()
try:
document = enhance(
args.caption,
args.base_url,
args.model,
api_key=args.api_key,
timeout=args.timeout,
temperature=args.temperature,
max_tokens=args.max_tokens,
)
except PromptEnhancementError as error:
print(f"pe_ling: {error}", file=sys.stderr)
raise SystemExit(1)
elapsed = time.perf_counter() - started
layer_count = len(document["layers"])
payload = json.dumps(document, indent=2, ensure_ascii=False) + "\n"
if args.out is not None:
args.out.parent.mkdir(parents=True, exist_ok=True)
args.out.write_text(payload, encoding="utf-8")
print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s) -> {args.out}")
else:
sys.stdout.write(payload)
print(f"pe_ling: {elapsed:.1f}s, {layer_count} layer(s)", file=sys.stderr)
if __name__ == "__main__":
main()