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,