[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:
co-authored by
Elizaveta Martirosian
Elizaveta Martirosian
ronnie_zheng
parent
87dc211b87
commit
44e4999ab2
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user