[diffusion] fix: respect --prompt-path (#21756)

This commit is contained in:
Mick
2026-04-01 16:47:59 +08:00
committed by GitHub
parent 95b881452e
commit 7bba319f1e
3 changed files with 111 additions and 4 deletions
@@ -183,7 +183,10 @@ class DiffGenerator:
multiple prompts, or None when every request failed.
"""
# 1. prepare requests
prompts = self._resolve_prompts(sampling_params_kwargs.get("prompt"))
prompts = self._resolve_prompts(
sampling_params_kwargs.get("prompt"),
sampling_params_kwargs.get("prompt_path"),
)
user_output_file_name = sampling_params_kwargs.get("output_file_name")
if len(prompts) > 1 and user_output_file_name is not None:
@@ -334,10 +337,14 @@ class DiffGenerator:
return None
return results[0] if len(results) == 1 else results
def _resolve_prompts(self, prompt: str | list[str] | None) -> list[str]:
def _resolve_prompts(
self,
prompt: str | list[str] | None,
prompt_path: str | None = None,
) -> list[str]:
"""Collect prompts from the argument or from a prompt file."""
if self.server_args.prompt_file_path is not None:
path = self.server_args.prompt_file_path
path = prompt_path or self.server_args.prompt_file_path
if path is not None:
if not os.path.exists(path):
raise FileNotFoundError(f"Prompt text file not found: {path}")
with open(path, encoding="utf-8") as f:
@@ -36,6 +36,7 @@ SUITES = {
"../unit/test_lora_format_adapter.py",
"../unit/test_server_args.py",
"../unit/test_input_validation.py",
"../unit/test_resolve_prompts.py",
# add new unit tests here
],
"1-gpu": [
@@ -0,0 +1,99 @@
import os
import tempfile
import unittest
from types import SimpleNamespace
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator
def _make_generator(prompt_file_path=None):
"""Return a DiffGenerator-shaped object with only server_args populated."""
obj = object.__new__(DiffGenerator)
obj.server_args = SimpleNamespace(prompt_file_path=prompt_file_path)
return obj
class TestResolvePrompts(unittest.TestCase):
# ---- inline prompt ----
def test_none_prompt_returns_space(self):
gen = _make_generator()
self.assertEqual(gen._resolve_prompts(None), [" "])
def test_string_prompt(self):
gen = _make_generator()
self.assertEqual(gen._resolve_prompts("hello"), ["hello"])
def test_list_prompt(self):
gen = _make_generator()
self.assertEqual(gen._resolve_prompts(["a", "b"]), ["a", "b"])
# ---- prompt_path (SamplingParams) ----
def test_prompt_path_single_line(self):
gen = _make_generator()
with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as f:
f.write("sunset over the ocean\n")
path = f.name
try:
result = gen._resolve_prompts(None, prompt_path=path)
self.assertEqual(result, ["sunset over the ocean"])
finally:
os.unlink(path)
def test_prompt_path_multi_line(self):
gen = _make_generator()
with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as f:
f.write("line one\n\nline two\n")
path = f.name
try:
result = gen._resolve_prompts(None, prompt_path=path)
self.assertEqual(result, ["line one", "line two"])
finally:
os.unlink(path)
def test_prompt_path_takes_priority_over_server_args(self):
with tempfile.NamedTemporaryFile(
"w", suffix=".txt", delete=False
) as f1, tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as f2:
f1.write("from prompt_path\n")
f2.write("from server_args\n")
path1, path2 = f1.name, f2.name
try:
gen = _make_generator(prompt_file_path=path2)
result = gen._resolve_prompts(None, prompt_path=path1)
self.assertEqual(result, ["from prompt_path"])
finally:
os.unlink(path1)
os.unlink(path2)
# ---- prompt_file_path (ServerArgs) ----
def test_server_args_prompt_file_path(self):
with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as f:
f.write("from server args\n")
path = f.name
try:
gen = _make_generator(prompt_file_path=path)
result = gen._resolve_prompts(None)
self.assertEqual(result, ["from server args"])
finally:
os.unlink(path)
# ---- error cases ----
def test_missing_file_raises(self):
gen = _make_generator()
with self.assertRaises(FileNotFoundError):
gen._resolve_prompts(None, prompt_path="/nonexistent/file.txt")
def test_empty_file_raises(self):
with tempfile.NamedTemporaryFile("w", suffix=".txt", delete=False) as f:
f.write(" \n\n \n")
path = f.name
try:
gen = _make_generator()
with self.assertRaises(ValueError):
gen._resolve_prompts(None, prompt_path=path)
finally:
os.unlink(path)
if __name__ == "__main__":
unittest.main()