[Diffusion] Use SGLang server for ERNIE-Image prompt enhancement (#31354)

Co-authored-by: Elizaveta Martirosian <you@example.com>
Co-authored-by: Elizaveta Martirosian <elizaveta.martirosian@gmail.com>
Co-authored-by: ronnie_zheng <zl19940307@163.com>
This commit is contained in:
Elizaveta Martirosian
2026-07-18 06:55:28 +03:00
committed by GitHub
co-authored by Elizaveta Martirosian Elizaveta Martirosian ronnie_zheng
parent 87dc211b87
commit 44e4999ab2
4 changed files with 105 additions and 0 deletions
@@ -2,6 +2,7 @@
import json
import os
import requests
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
@@ -79,6 +80,31 @@ class PEModelWrapper:
return self
class SGLangPEModelWrapper:
def __init__(self, model_url):
self.model_url = model_url.rstrip("/")
# Tokenizer is initialized separately during pipeline setup
self.pe_tokenizer = None
def generate(self, prompt: str, sampling_params: dict) -> dict:
response = requests.post(
self.model_url + "/generate",
json={
"text": prompt,
"sampling_params": sampling_params,
},
)
response.raise_for_status()
return response.json()
def to(self, *args, **kwargs):
logger.debug("Ignoring .to() because PE model is served externally")
return self
class PELoader(ComponentLoader):
"""Loader for prompt-enhancement causal LM (Ministral-3 based)."""
@@ -88,6 +114,12 @@ class PELoader(ComponentLoader):
def load_customized(
self, component_model_path: str, server_args: ServerArgs, component_name: str
):
if server_args.pe_server_url is not None:
logger.info(
f"Using external SGLang server for PE: {server_args.pe_server_url}"
)
return SGLangPEModelWrapper(server_args.pe_server_url)
logger.info("Loading PE model from %s ...", component_model_path)
pe_tokenizer_dir = os.path.join(
@@ -426,6 +426,9 @@ class ServerArgs(DisaggServerArgsMixin):
srt_encoder_connect_timeout: int = 3.05
srt_encoder_timeout: int = 100
# SGLang server for PE model inference
pe_server_url: str | None = None
@property
def broker_port(self) -> int:
return self.port + 1
@@ -1938,6 +1941,14 @@ class ServerArgs(DisaggServerArgsMixin):
"Increase value if connection between diffusion server and AR model server is slow.",
)
# SGLang server for PE model inference
parser.add_argument(
"--pe-server-url",
type=str,
default=ServerArgs.pe_server_url,
help="URL of SGLang server for PE model",
)
return parser
def url(self):