diff --git a/backend/python/diffusers/backend.py b/backend/python/diffusers/backend.py index 539ce54448bf..6020cd0e9f15 100755 --- a/backend/python/diffusers/backend.py +++ b/backend/python/diffusers/backend.py @@ -868,7 +868,7 @@ def GenerateImage(self, request, context): else: # pass the kwargs dictionary to the self.pipe method image = self.pipe( - prompt, + prompt=prompt, guidance_scale=self.cfg_scale, **kwargs ).images[0] diff --git a/backend/python/diffusers/test.py b/backend/python/diffusers/test.py index eff293ee6e10..2922f2a03f0b 100644 --- a/backend/python/diffusers/test.py +++ b/backend/python/diffusers/test.py @@ -373,3 +373,55 @@ def test_options_merged_into_pipeline_kwargs(self): finally: os.unlink(src_file.name) os.unlink(dst_file.name) + + def test_text_to_image_prompt_is_passed_by_keyword(self): + """Test compatibility with pipelines that take image before prompt.""" + import os + import tempfile + + from PIL import Image + + from backend import BackendServicer + + class Flux2CompatiblePipeline: + """Model the FLUX.2 call signature: image is before prompt.""" + + def __call__(self, image=None, prompt=None, **kwargs): + if prompt is None: + raise ValueError("prompt was not passed by keyword") + self.prompt = prompt + self.kwargs = kwargs + return MagicMock(images=[Image.new("RGB", (4, 4))]) + + pipeline = Flux2CompatiblePipeline() + svc = BackendServicer.__new__(BackendServicer) + svc.pipe = pipeline + svc.cfg_scale = 7.5 + svc.controlnet = None + svc.img2vid = False + svc.txt2vid = False + svc.clip_skip = 0 + svc.PipelineType = "Flux2KleinPipeline" + svc.options = {} + + with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as dst_file: + dst_path = dst_file.name + + try: + request = MagicMock() + request.positive_prompt = "a red apple on a wooden table" + request.negative_prompt = "" + request.step = 4 + request.seed = 0 + request.width = 0 + request.height = 0 + request.src = "" + request.ref_images = [] + request.dst = dst_path + + svc.GenerateImage(request, context=None) + + self.assertEqual(pipeline.prompt, request.positive_prompt) + self.assertEqual(pipeline.kwargs["num_inference_steps"], 4) + finally: + os.unlink(dst_path)