diff --git a/dashboard/src/lib/components/ImageParamsPanel.svelte b/dashboard/src/lib/components/ImageParamsPanel.svelte index a7d23798..ef17c278 100644 --- a/dashboard/src/lib/components/ImageParamsPanel.svelte +++ b/dashboard/src/lib/components/ImageParamsPanel.svelte @@ -110,6 +110,36 @@ setImageGenerationParams({ negativePrompt: value || null }); } + function handleNumImagesChange(event: Event) { + const input = event.target as HTMLInputElement; + const value = input.value.trim(); + if (value === "") { + setImageGenerationParams({ numImages: 1 }); + } else { + const num = parseInt(value, 10); + if (!isNaN(num) && num >= 1) { + setImageGenerationParams({ numImages: num }); + } + } + } + + function handleStreamChange(enabled: boolean) { + setImageGenerationParams({ stream: enabled }); + } + + function handlePartialImagesChange(event: Event) { + const input = event.target as HTMLInputElement; + const value = input.value.trim(); + if (value === "") { + setImageGenerationParams({ partialImages: 0 }); + } else { + const num = parseInt(value, 10); + if (!isNaN(num) && num >= 0) { + setImageGenerationParams({ partialImages: num }); + } + } + } + function clearSteps() { setImageGenerationParams({ numInferenceSteps: null }); } @@ -325,6 +355,59 @@ + + {#if !isEditMode} +
+ IMAGES: + +
+ {/if} + + +
+ STREAM: + +
+ + + {#if params.stream} +
+ PARTIALS: + +
+ {/if} + {#if isEditMode}
diff --git a/dashboard/src/lib/stores/app.svelte.ts b/dashboard/src/lib/stores/app.svelte.ts index d1aff1a1..ebfe0cb8 100644 --- a/dashboard/src/lib/stores/app.svelte.ts +++ b/dashboard/src/lib/stores/app.svelte.ts @@ -238,6 +238,10 @@ export interface ImageGenerationParams { size: "512x512" | "768x768" | "1024x1024" | "1024x768" | "768x1024"; quality: "low" | "medium" | "high"; outputFormat: "png" | "jpeg"; + numImages: number; + // Streaming params + stream: boolean; + partialImages: number; // Advanced params seed: number | null; numInferenceSteps: number | null; @@ -257,6 +261,9 @@ const DEFAULT_IMAGE_PARAMS: ImageGenerationParams = { size: "1024x1024", quality: "medium", outputFormat: "png", + numImages: 1, + stream: true, + partialImages: 3, seed: null, numInferenceSteps: null, guidance: null, @@ -1809,12 +1816,13 @@ class AppStore { const requestBody: Record = { model, prompt, + n: params.numImages, quality: params.quality, size: params.size, output_format: params.outputFormat, response_format: "b64_json", - stream: true, - partial_images: 3, + stream: params.stream, + partial_images: params.partialImages, }; if (hasAdvancedParams) { @@ -1878,31 +1886,74 @@ class AppStore { if (imageData && idx !== -1) { const format = parsed.format || "png"; const mimeType = `image/${format}`; + const imageIndex = parsed.image_index ?? 0; + const numImages = params.numImages; + if (parsed.type === "partial") { // Update with partial image and progress const partialNum = (parsed.partial_index ?? 0) + 1; const totalPartials = parsed.total_partials ?? 3; - this.messages[idx].content = - `Generating... ${partialNum}/${totalPartials}`; - this.messages[idx].attachments = [ - { - type: "generated-image", - name: `generated-image.${format}`, - preview: `data:${mimeType};base64,${imageData}`, - mimeType, - }, - ]; + const progressText = + numImages > 1 + ? `Generating image ${imageIndex + 1}/${numImages}... ${partialNum}/${totalPartials}` + : `Generating... ${partialNum}/${totalPartials}`; + this.messages[idx].content = progressText; + + const partialAttachment: MessageAttachment = { + type: "generated-image", + name: `generated-image.${format}`, + preview: `data:${mimeType};base64,${imageData}`, + mimeType, + }; + + if (imageIndex === 0) { + // First image - safe to replace attachments with partial preview + this.messages[idx].attachments = [partialAttachment]; + } else { + // Subsequent images - keep existing finals, show partial at current position + const existingAttachments = + this.messages[idx].attachments || []; + // Keep only the completed final images (up to current imageIndex) + const finals = existingAttachments.slice(0, imageIndex); + this.messages[idx].attachments = [ + ...finals, + partialAttachment, + ]; + } } else if (parsed.type === "final") { - // Final image - this.messages[idx].content = ""; - this.messages[idx].attachments = [ - { - type: "generated-image", - name: `generated-image.${format}`, - preview: `data:${mimeType};base64,${imageData}`, - mimeType, - }, - ]; + // Final image - replace partial at this position + const newAttachment: MessageAttachment = { + type: "generated-image", + name: `generated-image-${imageIndex + 1}.${format}`, + preview: `data:${mimeType};base64,${imageData}`, + mimeType, + }; + + if (imageIndex === 0) { + // First final image - replace any partial preview + this.messages[idx].attachments = [newAttachment]; + } else { + // Subsequent images - keep previous finals, replace partial at current position + const existingAttachments = + this.messages[idx].attachments || []; + // Slice keeps indices 0 to imageIndex-1 (the previous final images) + const previousFinals = existingAttachments.slice( + 0, + imageIndex, + ); + this.messages[idx].attachments = [ + ...previousFinals, + newAttachment, + ]; + } + + // Update progress message for multiple images + if (numImages > 1 && imageIndex < numImages - 1) { + this.messages[idx].content = + `Generating image ${imageIndex + 2}/${numImages}...`; + } else { + this.messages[idx].content = ""; + } } } } catch { @@ -1983,8 +2034,8 @@ class AppStore { formData.append("size", params.size); formData.append("output_format", params.outputFormat); formData.append("response_format", "b64_json"); - formData.append("stream", "1"); // Use "1" instead of "true" for reliable FastAPI boolean parsing - formData.append("partial_images", "3"); + formData.append("stream", params.stream ? "1" : "0"); + formData.append("partial_images", params.partialImages.toString()); formData.append("input_fidelity", params.inputFidelity); // Advanced params diff --git a/src/exo/master/api.py b/src/exo/master/api.py index a436f53e..f893fedf 100644 --- a/src/exo/master/api.py +++ b/src/exo/master/api.py @@ -835,6 +835,7 @@ class API: # Yield partial image event (always use b64_json for partials) event_data = { "type": "partial", + "image_index": chunk.image_index, "partial_index": partial_idx, "total_partials": total_partials, "format": str(chunk.format), diff --git a/src/exo/shared/types/worker/runner_response.py b/src/exo/shared/types/worker/runner_response.py index 8d695ab0..9f867413 100644 --- a/src/exo/shared/types/worker/runner_response.py +++ b/src/exo/shared/types/worker/runner_response.py @@ -30,6 +30,7 @@ class ImageGenerationResponse(BaseRunnerResponse): image_data: bytes format: Literal["png", "jpeg", "webp"] = "png" stats: ImageGenerationStats | None = None + image_index: int = 0 def __repr_args__(self) -> Generator[tuple[str, Any], None, None]: for name, value in super().__repr_args__(): # pyright: ignore[reportAny] @@ -44,6 +45,7 @@ class PartialImageResponse(BaseRunnerResponse): format: Literal["png", "jpeg", "webp"] = "png" partial_index: int total_partials: int + image_index: int = 0 def __repr_args__(self) -> Generator[tuple[str, Any], None, None]: for name, value in super().__repr_args__(): # pyright: ignore[reportAny] diff --git a/src/exo/worker/engines/image/generate.py b/src/exo/worker/engines/image/generate.py index 8bd34749..66f21a1e 100644 --- a/src/exo/worker/engines/image/generate.py +++ b/src/exo/worker/engines/image/generate.py @@ -75,19 +75,20 @@ def generate_image( intermediate images, then ImageGenerationResponse for the final image. Yields: - PartialImageResponse for intermediate images (if partial_images > 0) - ImageGenerationResponse for the final complete image + PartialImageResponse for intermediate images (if partial_images > 0, first image only) + ImageGenerationResponse for final complete images """ width, height = parse_size(task.size) quality: Literal["low", "medium", "high"] = task.quality or "medium" advanced_params = task.advanced_params if advanced_params is not None and advanced_params.seed is not None: - seed = advanced_params.seed + base_seed = advanced_params.seed else: - seed = random.randint(0, 2**32 - 1) + base_seed = random.randint(0, 2**32 - 1) is_bench = getattr(task, "bench", False) + num_images = task.n or 1 generation_start_time: float = 0.0 @@ -95,7 +96,11 @@ def generate_image( mx.reset_peak_memory() generation_start_time = time.perf_counter() - partial_images = task.partial_images or (3 if task.stream else 0) + partial_images = ( + task.partial_images + if task.partial_images is not None + else (3 if task.stream else 0) + ) image_path: Path | None = None @@ -105,72 +110,81 @@ def generate_image( image_path = Path(tmpdir) / "input.png" image_path.write_bytes(base64.b64decode(task.image_data)) - # Iterate over generator results - for result in model.generate( - prompt=task.prompt, - height=height, - width=width, - quality=quality, - seed=seed, - image_path=image_path, - partial_images=partial_images, - advanced_params=advanced_params, - ): - if isinstance(result, tuple): - # Partial image: (Image, partial_index, total_partials) - image, partial_idx, total_partials = result - buffer = io.BytesIO() - image_format = task.output_format.upper() - if image_format == "JPG": - image_format = "JPEG" - if image_format == "JPEG" and image.mode == "RGBA": - image = image.convert("RGB") - image.save(buffer, format=image_format) + for image_num in range(num_images): + # Increment seed for each image to ensure unique results + current_seed = base_seed + image_num - yield PartialImageResponse( - image_data=buffer.getvalue(), - format=task.output_format, - partial_index=partial_idx, - total_partials=total_partials, - ) - else: - image = result + for result in model.generate( + prompt=task.prompt, + height=height, + width=width, + quality=quality, + seed=current_seed, + image_path=image_path, + partial_images=partial_images, + advanced_params=advanced_params, + ): + if isinstance(result, tuple): + # Partial image: (Image, partial_index, total_partials) + image, partial_idx, total_partials = result + buffer = io.BytesIO() + image_format = task.output_format.upper() + if image_format == "JPG": + image_format = "JPEG" + if image_format == "JPEG" and image.mode == "RGBA": + image = image.convert("RGB") + image.save(buffer, format=image_format) - stats: ImageGenerationStats | None = None - if is_bench: - generation_end_time = time.perf_counter() - total_generation_time = generation_end_time - generation_start_time - - num_inference_steps = model.get_steps_for_quality(quality) - - seconds_per_step = ( - total_generation_time / num_inference_steps - if num_inference_steps > 0 - else 0.0 + yield PartialImageResponse( + image_data=buffer.getvalue(), + format=task.output_format, + partial_index=partial_idx, + total_partials=total_partials, + image_index=image_num, ) + else: + image = result - peak_memory_gb = mx.get_peak_memory() / (1024**3) + # Only include stats on the final image + stats: ImageGenerationStats | None = None + if is_bench and image_num == num_images - 1: + generation_end_time = time.perf_counter() + total_generation_time = ( + generation_end_time - generation_start_time + ) - stats = ImageGenerationStats( - seconds_per_step=seconds_per_step, - total_generation_time=total_generation_time, - num_inference_steps=num_inference_steps, - num_images=task.n or 1, - image_width=width, - image_height=height, - peak_memory_usage=Memory.from_gb(peak_memory_gb), + num_inference_steps = model.get_steps_for_quality(quality) + total_steps = num_inference_steps * num_images + + seconds_per_step = ( + total_generation_time / total_steps + if total_steps > 0 + else 0.0 + ) + + peak_memory_gb = mx.get_peak_memory() / (1024**3) + + stats = ImageGenerationStats( + seconds_per_step=seconds_per_step, + total_generation_time=total_generation_time, + num_inference_steps=num_inference_steps, + num_images=num_images, + image_width=width, + image_height=height, + peak_memory_usage=Memory.from_gb(peak_memory_gb), + ) + + buffer = io.BytesIO() + image_format = task.output_format.upper() + if image_format == "JPG": + image_format = "JPEG" + if image_format == "JPEG" and image.mode == "RGBA": + image = image.convert("RGB") + image.save(buffer, format=image_format) + + yield ImageGenerationResponse( + image_data=buffer.getvalue(), + format=task.output_format, + stats=stats, + image_index=image_num, ) - - buffer = io.BytesIO() - image_format = task.output_format.upper() - if image_format == "JPG": - image_format = "JPEG" - if image_format == "JPEG" and image.mode == "RGBA": - image = image.convert("RGB") - image.save(buffer, format=image_format) - - yield ImageGenerationResponse( - image_data=buffer.getvalue(), - format=task.output_format, - stats=stats, - ) diff --git a/src/exo/worker/runner/runner.py b/src/exo/worker/runner/runner.py index 24e42691..4f91deda 100644 --- a/src/exo/worker/runner/runner.py +++ b/src/exo/worker/runner/runner.py @@ -612,7 +612,7 @@ def _process_image_response( command_id=command_id, model_id=shard_metadata.model_card.model_id, event_sender=event_sender, - image_index=response.partial_index if is_partial else image_index, + image_index=response.image_index, is_partial=is_partial, partial_index=response.partial_index if is_partial else None, total_partials=response.total_partials if is_partial else None,