Add documentation for Image Generation

- Refine StabilityAI classes
- Add documentation
This commit is contained in:
Mark Pollack
2024-02-14 08:59:36 -05:00
parent 36278afdb3
commit 43dbaf4edd
17 changed files with 1007 additions and 323 deletions

View File

@@ -21,6 +21,8 @@ import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.image.ImageOptions;
import org.springframework.ai.openai.api.OpenAiImageApi;
import java.util.Objects;
/**
* OpenAI Image API options. OpenAiImageOptions.java
*
@@ -44,6 +46,18 @@ public class OpenAiImageOptions implements ImageOptions {
@JsonProperty("model")
private String model = OpenAiImageApi.DEFAULT_IMAGE_MODEL;
/**
* The width of the generated images. Must be one of 256, 512, or 1024 for dall-e-2.
*/
@JsonProperty("size_width")
private Integer width;
/**
* The height of the generated images. Must be one of 256, 512, or 1024 for dall-e-2.
*/
@JsonProperty("size_height")
private Integer height;
/**
* The quality of the image that will be generated. hd creates images with finer
* details and greater consistency across the image. This param is only supported for
@@ -66,18 +80,6 @@ public class OpenAiImageOptions implements ImageOptions {
@JsonProperty("size")
private String size;
/**
* The width of the generated images. Must be one of 256, 512, or 1024 for dall-e-2.
*/
@JsonProperty("size_width")
private Integer width;
/**
* The height of the generated images. Must be one of 256, 512, or 1024 for dall-e-2.
*/
@JsonProperty("size_height")
private Integer height;
/**
* The style of the generated images. Must be one of vivid or natural. Vivid causes
* the model to lean towards generating hyper-real and dramatic images. Natural causes
@@ -235,4 +237,28 @@ public class OpenAiImageOptions implements ImageOptions {
return (this.width != null && this.height != null) ? this.width + "x" + this.height : null;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof OpenAiImageOptions that))
return false;
return Objects.equals(n, that.n) && Objects.equals(model, that.model) && Objects.equals(width, that.width)
&& Objects.equals(height, that.height) && Objects.equals(quality, that.quality)
&& Objects.equals(responseFormat, that.responseFormat) && Objects.equals(size, that.size)
&& Objects.equals(style, that.style) && Objects.equals(user, that.user);
}
@Override
public int hashCode() {
return Objects.hash(n, model, width, height, quality, responseFormat, size, style, user);
}
@Override
public String toString() {
return "OpenAiImageOptions{" + "n=" + n + ", model='" + model + '\'' + ", width=" + width + ", height=" + height
+ ", quality='" + quality + '\'' + ", responseFormat='" + responseFormat + '\'' + ", size='" + size
+ '\'' + ", style='" + style + '\'' + ", user='" + user + '\'' + '}';
}
}

View File

@@ -21,8 +21,6 @@ import org.springframework.ai.image.*;
import org.springframework.ai.model.ModelOptionsUtils;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptions;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptionsBuilder;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptionsImpl;
import org.springframework.util.Assert;
import java.util.List;
@@ -41,7 +39,7 @@ public class StabilityAiImageClient implements ImageClient {
private final StabilityAiApi stabilityAiApi;
public StabilityAiImageClient(StabilityAiApi stabilityAiApi) {
this(stabilityAiApi, StabilityAiImageOptionsBuilder.builder().build());
this(stabilityAiApi, StabilityAiImageOptions.builder().build());
}
public StabilityAiImageClient(StabilityAiApi stabilityAiApi, StabilityAiImageOptions options) {
@@ -71,7 +69,7 @@ public class StabilityAiImageClient implements ImageClient {
// options configured via Autoconfiguration.
// Runtime options overwrite StabilityAiImageClient options
StabilityAiImageOptions optionsToUse = ModelOptionsUtils.merge(runtimeOptions, this.options,
StabilityAiImageOptionsImpl.class);
StabilityAiImageOptions.class);
// Copy the org.springframework.ai.model derived ImagePrompt and ImageOptions data
// types to the data types used in StabilityAiApi
@@ -100,7 +98,7 @@ public class StabilityAiImageClient implements ImageClient {
.withCfgScale(optionsToUse.getCfgScale())
.withClipGuidancePreset(optionsToUse.getClipGuidancePreset())
.withSampler(optionsToUse.getSampler())
.withSamples(optionsToUse.getSamples())
.withSamples(optionsToUse.getN())
.withSeed(optionsToUse.getSeed())
.withSteps(optionsToUse.getSteps())
.withStylePreset(optionsToUse.getStylePreset())
@@ -118,7 +116,7 @@ public class StabilityAiImageClient implements ImageClient {
}
private StabilityAiImageOptions convertOptions(ImageOptions runtimeOptions) {
StabilityAiImageOptionsBuilder builder = StabilityAiImageOptionsBuilder.builder();
StabilityAiImageOptions.Builder builder = StabilityAiImageOptions.builder();
if (runtimeOptions == null) {
return builder.build();
}

View File

@@ -15,28 +15,476 @@
*/
package org.springframework.ai.stabilityai.api;
import com.fasterxml.jackson.annotation.JsonInclude;
import com.fasterxml.jackson.annotation.JsonProperty;
import org.springframework.ai.image.ImageOptions;
import org.springframework.ai.stabilityai.StyleEnum;
import java.util.Objects;
/**
* StabilityAiImageOptions is an interface that extends ImageOptions. It provides
* additional stability AI specific image options.
*/
public interface StabilityAiImageOptions extends ImageOptions {
@JsonInclude(JsonInclude.Include.NON_NULL)
public class StabilityAiImageOptions implements ImageOptions {
Float getCfgScale();
/**
* The number of images to be generated.
*
* Defaults to 1 if not explicitly set, indicating a single image will be generated.
*
* <p>
* This method specifies the total number of images to generate. It allows for
* controlling the volume of output from a single operation, facilitating batch
* generation of images based on the provided settings.
* </p>
*
* <p>
* Valid range of values: 1 to 10. This ensures that the request remains within a
* manageable scale and aligns with system capabilities or limitations.
* </p>
*
*
*/
@JsonProperty("samples")
private Integer n;
String getClipGuidancePreset();
/**
* The engine/model to use in Stability AI The model is passed in the URL as a path
* parameter
*
* The default value is stable-diffusion-v1-6
*/
private String model = StabilityAiApi.DEFAULT_IMAGE_MODEL;
String getSampler();
/**
* Retrieves the width of the image to be generated, in pixels.
* <p>
* Specifies the desired width for the output image. The value must be a multiple of
* 64 and at least 128 pixels. This parameter is adjusted to comply with the
* specifications of the selected generation engine, which may have unique
* requirements based on its version.
* </p>
*
* <p>
* Default value: 512.
* </p>
*
* <p>
* Engine-specific dimension validation:
* </p>
* <ul>
* <li>SDXL Beta: Width must be between 128 and 896 pixels, with only one dimension
* allowed to exceed 512.</li>
* <li>SDXL v0.9 and v1.0: Width must match one of the predefined dimension
* pairs.</li>
* <li>SD v1.6: Width must be between 320 and 1536 pixels.</li>
* </ul>
*
*/
@JsonProperty("width")
private Integer width;
Integer getSamples();
/**
* Retrieves the height of the image to be generated, in pixels.
* <p>
* Specifies the desired height for the output image. The value must be a multiple of
* 64 and at least 128 pixels. This setting is crucial for ensuring compatibility with
* the underlying generation engine, which may impose additional restrictions based on
* the engine version.
* </p>
*
* <p>
* Default value: 512.
* </p>
*
* <p>
* Engine-specific dimension validation:
* </p>
* <ul>
* <li>SDXL Beta: Height must be between 128 and 896 pixels, with only one dimension
* allowed to exceed 512.</li>
* <li>SDXL v0.9 and v1.0: Height must match one of the predefined dimension
* pairs.</li>
* <li>SD v1.6: Height must be between 320 and 1536 pixels.</li>
* </ul>
*
*/
@JsonProperty("height")
private Integer height;
Long getSeed();
/**
* The format in which the generated images are returned. It is sent as part of the
* accept header. Must be "application/json" or "image/png"
*/
private String responseFormat;
Integer getSteps();
/**
* The strictness level of the diffusion process adherence to the prompt text.
* <p>
* This field determines how closely the generated image will match the provided
* prompt. Higher values indicate that the image will adhere more closely to the
* prompt text, ensuring a closer match to the expected output.
* </p>
*
* <ul>
* <li>Range: 0 to 35</li>
* <li>Default value: 7</li>
* </ul>
*
*/
@JsonProperty("cfg_scale")
private Float cfgScale;
String getStylePreset();
/**
* The preset for clip guidance.
* <p>
* This field indicates the preset configuration for clip guidance, affecting the
* processing speed and characteristics. The choice of preset can influence the
* behavior of the guidance system, potentially impacting performance and output
* quality.
* </p>
*
* <p>
* Available presets are:
* <ul>
* <li>{@code FAST_BLUE}: An optimized preset for quicker processing with a focus on
* blue tones.</li>
* <li>{@code FAST_GREEN}: An optimized preset for quicker processing with a focus on
* green tones.</li>
* <li>{@code NONE}: No preset is applied, default processing.</li>
* <li>{@code SIMPLE}: A basic level of clip guidance for general use.</li>
* <li>{@code SLOW}: A slower processing preset for more detailed guidance.</li>
* <li>{@code SLOWER}: Further reduces the processing speed for enhanced detail in
* guidance.</li>
* <li>{@code SLOWEST}: The slowest processing speed, offering the highest level of
* detail in clip guidance.</li>
* </ul>
* </p>
*
* Defaults to {@code NONE} if no specific preset is configured.
*
*/
@JsonProperty("clip_guidance_preset")
private String clipGuidancePreset;
// extras json object...
/**
* The name of the sampler used for the diffusion process.
* <p>
* This field specifies the sampler algorithm to be used during the diffusion process.
* Selecting a specific sampler can influence the quality and characteristics of the
* generated output. If no sampler is explicitly selected, an appropriate sampler will
* be automatically chosen based on the context or other settings.
* </p>
*
* <p>
* Available samplers are:
* <ul>
* <li>{@code DDIM}: A deterministic diffusion inverse model for stable and
* predictable outputs.</li>
* <li>{@code DDPM}: Denoising diffusion probabilistic models for high-quality
* generation.</li>
* <li>{@code K_DPMPP_2M}: A specific configuration of DPM++ model with medium
* settings.</li>
* <li>{@code K_DPMPP_2S_ANCESTRAL}: An ancestral sampling variant of the DPM++ model
* with small settings.</li>
* <li>{@code K_DPM_2}: A variant of the DPM model designed for balanced
* performance.</li>
* <li>{@code K_DPM_2_ANCESTRAL}: An ancestral sampling variant of the DPM model.</li>
* <li>{@code K_EULER}: Utilizes the Euler method for diffusion, offering a different
* trade-off between speed and quality.</li>
* <li>{@code K_EULER_ANCESTRAL}: An ancestral version of the Euler method for nuanced
* sampling control.</li>
* <li>{@code K_HEUN}: Employs the Heun's method for a more accurate approximation in
* the diffusion process.</li>
* <li>{@code K_LMS}: Leverages the linear multistep method for potentially improved
* diffusion quality.</li>
* </ul>
* </p>
*
* An appropriate sampler is automatically selected if this value is omitted.
*
*/
@JsonProperty("sampler")
private String sampler;
/**
* The seed used for generating random noise.
* <p>
* This value serves as the seed for random noise generation, influencing the
* randomness and uniqueness of the output. A specific seed ensures reproducibility of
* results. Omitting this option or using 0 triggers the selection of a random seed.
* </p>
*
* <p>
* Valid range of values: 0 to 4294967295.
* </p>
*
* Default is 0, which indicates that a random seed will be used.
*/
@JsonProperty("seed")
private Long seed;
/**
* The number of diffusion steps to run.
* <p>
* Specifies the total number of steps in the diffusion process, affecting the detail
* and quality of the generated output. More steps can lead to higher quality but
* require more processing time.
* </p>
*
* <p>
* Valid range of values: 10 to 50.
* </p>
*
* Defaults to 30 if not explicitly set.
*/
@JsonProperty("steps")
private Integer steps;
/**
* The style preset intended to guide the image model towards a specific artistic
* style.
* <p>
* This string parameter allows for the selection of a predefined style preset,
* influencing the aesthetic characteristics of the generated image. The choice of
* preset can significantly impact the visual outcome, aligning it with particular
* artistic genres or techniques.
* </p>
*
* <p>
* Possible values include:
* </p>
* <ul>
* <li>{@code 3d-model}</li>
* <li>{@code analog-film}</li>
* <li>{@code anime}</li>
* <li>{@code cinematic}</li>
* <li>{@code comic-book}</li>
* <li>{@code digital-art}</li>
* <li>{@code enhance}</li>
* <li>{@code fantasy-art}</li>
* <li>{@code isometric}</li>
* <li>{@code line-art}</li>
* <li>{@code low-poly}</li>
* <li>{@code modeling-compound}</li>
* <li>{@code neon-punk}</li>
* <li>{@code origami}</li>
* <li>{@code photographic}</li>
* <li>{@code pixel-art}</li>
* <li>{@code tile-texture}</li>
* </ul>
* <p>
* Note: This list of style presets is subject to change.
* </p>
*
*/
@JsonProperty("style_preset")
private String stylePreset;
public static Builder builder() {
return new Builder();
}
public static class Builder {
private final StabilityAiImageOptions options;
private Builder() {
this.options = new StabilityAiImageOptions();
}
public Builder withN(Integer n) {
options.setN(n);
return this;
}
public Builder withModel(String model) {
options.setModel(model);
return this;
}
public Builder withWidth(Integer width) {
options.setWidth(width);
return this;
}
public Builder withHeight(Integer height) {
options.setHeight(height);
return this;
}
public Builder withResponseFormat(String responseFormat) {
options.setResponseFormat(responseFormat);
return this;
}
public Builder withCfgScale(Float cfgScale) {
options.setCfgScale(cfgScale);
return this;
}
public Builder withClipGuidancePreset(String clipGuidancePreset) {
options.setClipGuidancePreset(clipGuidancePreset);
return this;
}
public Builder withSampler(String sampler) {
options.setSampler(sampler);
return this;
}
public Builder withSeed(Long seed) {
options.setSeed(seed);
return this;
}
public Builder withSteps(Integer steps) {
options.setSteps(steps);
return this;
}
public Builder withSamples(Integer samples) {
options.setN(samples);
return this;
}
public Builder withStylePreset(String stylePreset) {
options.setStylePreset(stylePreset);
return this;
}
public Builder withStylePreset(StyleEnum styleEnum) {
options.setStylePreset(styleEnum.toString());
return this;
}
public StabilityAiImageOptions build() {
return options;
}
}
@Override
public Integer getN() {
return n;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public String getModel() {
return model;
}
public void setModel(String model) {
this.model = model;
}
@Override
public Integer getWidth() {
return width;
}
public void setWidth(Integer width) {
this.width = width;
}
@Override
public Integer getHeight() {
return height;
}
public void setHeight(Integer height) {
this.height = height;
}
@Override
public String getResponseFormat() {
return responseFormat;
}
public void setResponseFormat(String responseFormat) {
this.responseFormat = responseFormat;
}
public Float getCfgScale() {
return cfgScale;
}
public void setCfgScale(Float cfgScale) {
this.cfgScale = cfgScale;
}
public String getClipGuidancePreset() {
return clipGuidancePreset;
}
public void setClipGuidancePreset(String clipGuidancePreset) {
this.clipGuidancePreset = clipGuidancePreset;
}
public String getSampler() {
return sampler;
}
public void setSampler(String sampler) {
this.sampler = sampler;
}
public Long getSeed() {
return seed;
}
public void setSeed(Long seed) {
this.seed = seed;
}
public Integer getSteps() {
return steps;
}
public void setSteps(Integer steps) {
this.steps = steps;
}
public String getStylePreset() {
return stylePreset;
}
public void setStylePreset(String stylePreset) {
this.stylePreset = stylePreset;
}
@Override
public boolean equals(Object o) {
if (this == o)
return true;
if (!(o instanceof StabilityAiImageOptions that))
return false;
return Objects.equals(n, that.n) && Objects.equals(model, that.model) && Objects.equals(width, that.width)
&& Objects.equals(height, that.height) && Objects.equals(responseFormat, that.responseFormat)
&& Objects.equals(cfgScale, that.cfgScale)
&& Objects.equals(clipGuidancePreset, that.clipGuidancePreset) && Objects.equals(sampler, that.sampler)
&& Objects.equals(seed, that.seed) && Objects.equals(steps, that.steps)
&& Objects.equals(stylePreset, that.stylePreset);
}
@Override
public int hashCode() {
return Objects.hash(n, model, width, height, responseFormat, cfgScale, clipGuidancePreset, sampler, seed, steps,
stylePreset);
}
@Override
public String toString() {
return "StabilityAiImageOptions{" + "n=" + n + ", model='" + model + '\'' + ", width=" + width + ", height="
+ height + ", responseFormat='" + responseFormat + '\'' + ", cfgScale=" + cfgScale
+ ", clipGuidancePreset='" + clipGuidancePreset + '\'' + ", sampler='" + sampler + '\'' + ", seed="
+ seed + ", steps=" + steps + ", stylePreset='" + stylePreset + '\'' + '}';
}
}

View File

@@ -1,106 +0,0 @@
/*
* Copyright 2023-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.stabilityai.api;
import org.springframework.ai.stabilityai.StyleEnum;
/**
* The StabilityAiImageOptionsBuilder class provides a convenient way to construct an
* instance of StabilityAiImageOptions. by allowing you to chain multiple method calls to
* set the desired options.
*/
public class StabilityAiImageOptionsBuilder {
private StabilityAiImageOptionsImpl options;
private StabilityAiImageOptionsBuilder() {
options = new StabilityAiImageOptionsImpl();
}
public static StabilityAiImageOptionsBuilder builder() {
return new StabilityAiImageOptionsBuilder();
}
public StabilityAiImageOptionsBuilder withN(Integer n) {
options.setN(n);
return this;
}
public StabilityAiImageOptionsBuilder withModel(String model) {
options.setModel(model);
return this;
}
public StabilityAiImageOptionsBuilder withWidth(Integer width) {
options.setWidth(width);
return this;
}
public StabilityAiImageOptionsBuilder withHeight(Integer height) {
options.setHeight(height);
return this;
}
public StabilityAiImageOptionsBuilder withResponseFormat(String responseFormat) {
options.setResponseFormat(responseFormat);
return this;
}
public StabilityAiImageOptionsBuilder withCfgScale(Float cfgScale) {
options.setCfgScale(cfgScale);
return this;
}
public StabilityAiImageOptionsBuilder withClipGuidancePreset(String clipGuidancePreset) {
options.setClipGuidancePreset(clipGuidancePreset);
return this;
}
public StabilityAiImageOptionsBuilder withSampler(String sampler) {
options.setSampler(sampler);
return this;
}
public StabilityAiImageOptionsBuilder withSeed(Long seed) {
options.setSeed(seed);
return this;
}
public StabilityAiImageOptionsBuilder withSteps(Integer steps) {
options.setSteps(steps);
return this;
}
public StabilityAiImageOptionsBuilder withSamples(Integer samples) {
options.setSamples(samples);
return this;
}
public StabilityAiImageOptionsBuilder withStylePreset(String stylePreset) {
options.setStylePreset(stylePreset);
return this;
}
public StabilityAiImageOptionsBuilder withStylePreset(StyleEnum styleEnum) {
options.setStylePreset(styleEnum.toString());
return this;
}
public StabilityAiImageOptions build() {
return options;
}
}

View File

@@ -1,155 +0,0 @@
/*
* Copyright 2023-2024 the original author or authors.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* https://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
package org.springframework.ai.stabilityai.api;
public class StabilityAiImageOptionsImpl implements StabilityAiImageOptions {
private Integer n;
private String model;
private Integer width;
private Integer height;
private String responseFormat;
private Float cfgScale;
private String clipGuidancePreset;
private String sampler;
private Integer samples;
private Long seed;
private Integer steps;
private String stylePreset;
public StabilityAiImageOptionsImpl() {
}
@Override
public Integer getN() {
return this.n;
}
public void setN(Integer n) {
this.n = n;
}
@Override
public String getModel() {
return this.model;
}
public void setModel(String model) {
this.model = model;
}
@Override
public Integer getWidth() {
return this.width;
}
public void setWidth(Integer width) {
this.width = width;
}
@Override
public Integer getHeight() {
return this.height;
}
public void setHeight(Integer height) {
this.height = height;
}
@Override
public String getResponseFormat() {
return this.responseFormat;
}
public void setResponseFormat(String responseFormat) {
this.responseFormat = responseFormat;
}
@Override
public Float getCfgScale() {
return this.cfgScale;
}
public void setCfgScale(Float cfgScale) {
this.cfgScale = cfgScale;
}
@Override
public String getClipGuidancePreset() {
return this.clipGuidancePreset;
}
public void setClipGuidancePreset(String clipGuidancePreset) {
this.clipGuidancePreset = clipGuidancePreset;
}
@Override
public String getSampler() {
return this.sampler;
}
public void setSampler(String sampler) {
this.sampler = sampler;
}
@Override
public Integer getSamples() {
return this.samples;
}
public void setSamples(Integer samples) {
this.samples = samples;
}
@Override
public Long getSeed() {
return this.seed;
}
public void setSeed(Long seed) {
this.seed = seed;
}
@Override
public Integer getSteps() {
return this.steps;
}
public void setSteps(Integer steps) {
this.steps = steps;
}
@Override
public String getStylePreset() {
return this.stylePreset;
}
public void setStylePreset(String stylePreset) {
this.stylePreset = stylePreset;
}
}

View File

@@ -18,14 +18,11 @@ package org.springframework.ai.stabilityai;
import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.condition.EnabledIfEnvironmentVariable;
import org.springframework.ai.image.*;
import org.springframework.ai.stabilityai.api.StabilityAiApi;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptions;
import org.springframework.ai.stabilityai.api.StabilityAiImageOptionsBuilder;
import org.springframework.beans.factory.annotation.Autowired;
import org.springframework.boot.test.context.SpringBootTest;
import java.io.File;
import java.io.FileNotFoundException;
import java.io.FileOutputStream;
import java.io.IOException;
import java.util.Base64;
@@ -42,12 +39,13 @@ public class StabilityAiImageClientIT {
@Test
void imageAsBase64Test() throws IOException {
StabilityAiImageOptions imageOptions = StabilityAiImageOptionsBuilder.builder()
StabilityAiImageOptions imageOptions = StabilityAiImageOptions.builder()
.withStylePreset(StyleEnum.PHOTOGRAPHIC)
.build();
var instructions = """
A light cream colored mini golden doodle with a sign that contains the message "I'm on my way to BARCADE!".""";
A light cream colored mini golden doodle.
""";
ImagePrompt imagePrompt = new ImagePrompt(instructions, imageOptions);
@@ -58,7 +56,7 @@ public class StabilityAiImageClientIT {
assertThat(image.getB64Json()).isNotEmpty();
// writeFile(image);
writeFile(image);
}
private static void writeFile(Image image) throws IOException {