mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-07-23 11:20:53 -05:00
refactor: return bool from image and upscale APIs (#1728)
This commit is contained in:
@@ -766,8 +766,12 @@ int main(int argc, const char* argv[]) {
|
||||
if (cli_params.mode == IMG_GEN) {
|
||||
sd_img_gen_params_t img_gen_params = gen_params.to_sd_img_gen_params_t();
|
||||
|
||||
num_results = gen_params.batch_count;
|
||||
results.adopt(generate_image(sd_ctx.get(), &img_gen_params), num_results);
|
||||
sd_image_t* generated_images = nullptr;
|
||||
if (!generate_image(sd_ctx.get(), &img_gen_params, &generated_images, &num_results)) {
|
||||
generated_images = nullptr;
|
||||
num_results = 0;
|
||||
}
|
||||
results.adopt(generated_images, num_results);
|
||||
} else if (cli_params.mode == VID_GEN) {
|
||||
sd_vid_gen_params_t vid_gen_params = gen_params.to_sd_vid_gen_params_t();
|
||||
sd_image_t* generated_video = nullptr;
|
||||
@@ -802,15 +806,21 @@ int main(int argc, const char* argv[]) {
|
||||
SDImageOwner current_image(results[i]);
|
||||
results[i] = {0, 0, 0, nullptr};
|
||||
for (int u = 0; u < gen_params.upscale_repeats; ++u) {
|
||||
sd_image_t* upscaled_images = upscale(upscaler_ctx.get(), current_image.get(), upscale_factor);
|
||||
if (upscaled_images == nullptr || upscaled_images[0].data == nullptr) {
|
||||
free_sd_images(upscaled_images, 1);
|
||||
sd_image_t* upscaled_images = nullptr;
|
||||
int upscaled_count = 0;
|
||||
bool upscale_ok = upscale(upscaler_ctx.get(),
|
||||
current_image.get(),
|
||||
upscale_factor,
|
||||
&upscaled_images,
|
||||
&upscaled_count);
|
||||
if (!upscale_ok || upscaled_count <= 0 || upscaled_images[0].data == nullptr) {
|
||||
free_sd_images(upscaled_images, upscaled_count);
|
||||
LOG_ERROR("upscale failed");
|
||||
break;
|
||||
}
|
||||
sd_image_t upscaled_image = upscaled_images[0];
|
||||
upscaled_images[0] = {0, 0, 0, nullptr};
|
||||
free_sd_images(upscaled_images, 1);
|
||||
free_sd_images(upscaled_images, upscaled_count);
|
||||
current_image.reset(upscaled_image);
|
||||
}
|
||||
results[i] = current_image.release(); // Set the final upscaled image as the result
|
||||
|
||||
@@ -173,8 +173,13 @@ bool execute_img_gen_job(ServerRuntime& runtime,
|
||||
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex);
|
||||
sd_image_t* raw_results = generate_image(runtime.sd_ctx, ¶ms);
|
||||
results.adopt(raw_results, params.batch_count);
|
||||
sd_image_t* raw_results = nullptr;
|
||||
int num_results = 0;
|
||||
if (!generate_image(runtime.sd_ctx, ¶ms, &raw_results, &num_results)) {
|
||||
raw_results = nullptr;
|
||||
num_results = 0;
|
||||
}
|
||||
results.adopt(raw_results, num_results);
|
||||
}
|
||||
|
||||
const int num_results = results.count();
|
||||
|
||||
@@ -229,8 +229,11 @@ static bool execute_sync_img_gen_request(ServerRuntime& runtime,
|
||||
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex);
|
||||
sd_image_t* raw_results = generate_image(runtime.sd_ctx, &img_gen_params);
|
||||
num_results = request.gen_params.batch_count;
|
||||
sd_image_t* raw_results = nullptr;
|
||||
if (!generate_image(runtime.sd_ctx, &img_gen_params, &raw_results, &num_results)) {
|
||||
raw_results = nullptr;
|
||||
num_results = 0;
|
||||
}
|
||||
results.adopt(raw_results, num_results);
|
||||
}
|
||||
|
||||
|
||||
@@ -292,8 +292,11 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {
|
||||
|
||||
{
|
||||
std::lock_guard<std::mutex> lock(*runtime->sd_ctx_mutex);
|
||||
sd_image_t* raw_results = generate_image(runtime->sd_ctx, &img_gen_params);
|
||||
num_results = request.gen_params.batch_count;
|
||||
sd_image_t* raw_results = nullptr;
|
||||
if (!generate_image(runtime->sd_ctx, &img_gen_params, &raw_results, &num_results)) {
|
||||
raw_results = nullptr;
|
||||
num_results = 0;
|
||||
}
|
||||
results.adopt(raw_results, num_results);
|
||||
}
|
||||
|
||||
|
||||
@@ -454,7 +454,10 @@ SD_API enum scheduler_t sd_get_default_scheduler(const sd_ctx_t* sd_ctx, enum sa
|
||||
|
||||
SD_API void sd_img_gen_params_init(sd_img_gen_params_t* sd_img_gen_params);
|
||||
SD_API char* sd_img_gen_params_to_str(const sd_img_gen_params_t* sd_img_gen_params);
|
||||
SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* sd_img_gen_params);
|
||||
SD_API bool generate_image(sd_ctx_t* sd_ctx,
|
||||
const sd_img_gen_params_t* sd_img_gen_params,
|
||||
sd_image_t** images_out,
|
||||
int* num_images_out);
|
||||
|
||||
enum sd_cancel_mode_t {
|
||||
// Stop the current generation as soon as possible.
|
||||
@@ -484,9 +487,11 @@ SD_API upscaler_ctx_t* new_upscaler_ctx(const char* esrgan_path,
|
||||
const char* params_backend);
|
||||
SD_API void free_upscaler_ctx(upscaler_ctx_t* upscaler_ctx);
|
||||
|
||||
SD_API sd_image_t* upscale(upscaler_ctx_t* upscaler_ctx,
|
||||
sd_image_t input_image,
|
||||
uint32_t upscale_factor);
|
||||
SD_API bool upscale(upscaler_ctx_t* upscaler_ctx,
|
||||
sd_image_t input_image,
|
||||
uint32_t upscale_factor,
|
||||
sd_image_t** images_out,
|
||||
int* num_images_out);
|
||||
|
||||
SD_API int get_upscale_factor(upscaler_ctx_t* upscaler_ctx);
|
||||
|
||||
|
||||
@@ -4278,7 +4278,8 @@ static std::optional<ImageGenerationEmbeds> prepare_image_generation_embeds(sd_c
|
||||
|
||||
static sd_image_t* decode_image_outputs(sd_ctx_t* sd_ctx,
|
||||
const GenerationRequest& request,
|
||||
const std::vector<sd::Tensor<float>>& final_latents) {
|
||||
const std::vector<sd::Tensor<float>>& final_latents,
|
||||
int* num_images_out) {
|
||||
if (final_latents.empty()) {
|
||||
LOG_ERROR("no latent images to decode");
|
||||
return nullptr;
|
||||
@@ -4320,11 +4321,14 @@ static sd_image_t* decode_image_outputs(sd_ctx_t* sd_ctx,
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
sd_image_t* result_images = (sd_image_t*)calloc(request.batch_count, sizeof(sd_image_t));
|
||||
int image_count = static_cast<int>(decoded_images.size());
|
||||
sd_image_t* result_images = (sd_image_t*)calloc(image_count, sizeof(sd_image_t));
|
||||
if (result_images == nullptr) {
|
||||
return nullptr;
|
||||
}
|
||||
memset(result_images, 0, request.batch_count * sizeof(sd_image_t));
|
||||
if (num_images_out != nullptr) {
|
||||
*num_images_out = image_count;
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < decoded_images.size(); i++) {
|
||||
result_images[i] = tensor_to_sd_image(decoded_images[i]);
|
||||
@@ -4517,9 +4521,18 @@ static std::vector<float> make_hires_sigma_schedule(sd_ctx_t* sd_ctx,
|
||||
sigmas.end());
|
||||
}
|
||||
|
||||
SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* sd_img_gen_params) {
|
||||
SD_API bool generate_image(sd_ctx_t* sd_ctx,
|
||||
const sd_img_gen_params_t* sd_img_gen_params,
|
||||
sd_image_t** images_out,
|
||||
int* num_images_out) {
|
||||
if (images_out != nullptr) {
|
||||
*images_out = nullptr;
|
||||
}
|
||||
if (num_images_out != nullptr) {
|
||||
*num_images_out = 0;
|
||||
}
|
||||
if (sd_ctx == nullptr || sd_img_gen_params == nullptr) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
sd_ctx->sd->reset_cancel_flag();
|
||||
@@ -4542,7 +4555,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
&request,
|
||||
&plan);
|
||||
if (!latents_opt.has_value()) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
ImageGenerationLatents latents = std::move(*latents_opt);
|
||||
|
||||
@@ -4552,7 +4565,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
&plan,
|
||||
&latents);
|
||||
if (!embeds_opt.has_value()) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
ImageGenerationEmbeds embeds = std::move(*embeds_opt);
|
||||
|
||||
@@ -4562,7 +4575,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
sd_cancel_mode_t cancel = sd_ctx->sd->get_cancel_flag();
|
||||
if (cancel == SD_CANCEL_ALL) {
|
||||
LOG_ERROR("cancelling generation");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
if (cancel == SD_CANCEL_NEW_LATENTS) {
|
||||
LOG_INFO("cancelling new latent generation, returning %zu/%d completed latents",
|
||||
@@ -4614,7 +4627,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
b + 1,
|
||||
request.batch_count,
|
||||
(sampling_end - sampling_start) * 1.0f / 1000);
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
int64_t denoise_end = ggml_time_ms();
|
||||
LOG_INFO("generating %zu latent images completed, taking %.2fs",
|
||||
@@ -4622,13 +4635,13 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
(denoise_end - denoise_start) * 1.0f / 1000);
|
||||
if (final_latents.empty()) {
|
||||
LOG_ERROR("no latent images generated");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
if (request.hires.enabled && request.hires.target_width > 0) {
|
||||
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
|
||||
LOG_ERROR("cancelling generation before hires fix");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
LOG_INFO("hires fix: upscaling to %dx%d", request.hires.target_width, request.hires.target_height);
|
||||
|
||||
@@ -4636,7 +4649,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
if (request.hires.upscaler == SD_HIRES_UPSCALER_MODEL) {
|
||||
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
|
||||
LOG_ERROR("cancelling generation before hires model load");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
LOG_INFO("hires fix: loading model upscaler from '%s'", request.hires.model_path);
|
||||
hires_upscaler = std::make_unique<UpscalerGGML>(sd_ctx->sd->n_threads,
|
||||
@@ -4649,7 +4662,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
if (!hires_upscaler->load_from_file(request.hires.model_path,
|
||||
sd_ctx->sd->n_threads)) {
|
||||
LOG_ERROR("load hires model upscaler failed");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4673,7 +4686,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
for (int b = 0; b < (int)final_latents.size(); b++) {
|
||||
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
|
||||
LOG_ERROR("cancelling generation during hires fix");
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
int64_t cur_seed = request.seed + b;
|
||||
sd_ctx->sd->rng->manual_seed(cur_seed);
|
||||
@@ -4684,7 +4697,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
request,
|
||||
hires_upscaler.get());
|
||||
if (upscaled.empty()) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
sd::Tensor<float> noise = sd::randn_like<float>(upscaled, sd_ctx->sd->rng);
|
||||
@@ -4738,7 +4751,7 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
b + 1,
|
||||
(int)final_latents.size(),
|
||||
(hires_sample_end - hires_sample_start) * 1.0f / 1000);
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
int64_t hires_denoise_end = ggml_time_ms();
|
||||
LOG_INFO("hires fix completed, taking %.2fs", (hires_denoise_end - hires_denoise_start) * 1.0f / 1000);
|
||||
@@ -4746,16 +4759,25 @@ SD_API sd_image_t* generate_image(sd_ctx_t* sd_ctx, const sd_img_gen_params_t* s
|
||||
final_latents = std::move(hires_final_latents);
|
||||
}
|
||||
|
||||
auto result = decode_image_outputs(sd_ctx, request, final_latents);
|
||||
int num_images = 0;
|
||||
auto result = decode_image_outputs(sd_ctx, request, final_latents, &num_images);
|
||||
if (result == nullptr) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
sd_ctx->sd->lora_stat();
|
||||
|
||||
int64_t t1 = ggml_time_ms();
|
||||
LOG_INFO("generate_image completed in %.2fs", (t1 - t0) * 1.0f / 1000);
|
||||
return result;
|
||||
if (num_images_out != nullptr) {
|
||||
*num_images_out = num_images;
|
||||
}
|
||||
if (images_out != nullptr) {
|
||||
*images_out = result;
|
||||
} else {
|
||||
free_sd_images(result, num_images);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
static std::optional<ImageGenerationLatents> prepare_video_generation_latents(sd_ctx_t* sd_ctx,
|
||||
|
||||
@@ -199,23 +199,41 @@ upscaler_ctx_t* new_upscaler_ctx(const char* esrgan_path_c_str,
|
||||
return upscaler_ctx;
|
||||
}
|
||||
|
||||
sd_image_t* upscale(upscaler_ctx_t* upscaler_ctx, sd_image_t input_image, uint32_t upscale_factor) {
|
||||
bool upscale(upscaler_ctx_t* upscaler_ctx,
|
||||
sd_image_t input_image,
|
||||
uint32_t upscale_factor,
|
||||
sd_image_t** images_out,
|
||||
int* num_images_out) {
|
||||
if (images_out != nullptr) {
|
||||
*images_out = nullptr;
|
||||
}
|
||||
if (num_images_out != nullptr) {
|
||||
*num_images_out = 0;
|
||||
}
|
||||
if (upscaler_ctx == nullptr || upscaler_ctx->upscaler == nullptr) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
sd_image_t* result_images = (sd_image_t*)calloc(1, sizeof(sd_image_t));
|
||||
if (result_images == nullptr) {
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
result_images[0] = upscaler_ctx->upscaler->upscale(input_image, upscale_factor);
|
||||
if (result_images[0].data == nullptr) {
|
||||
free(result_images);
|
||||
return nullptr;
|
||||
return false;
|
||||
}
|
||||
|
||||
return result_images;
|
||||
if (num_images_out != nullptr) {
|
||||
*num_images_out = 1;
|
||||
}
|
||||
if (images_out != nullptr) {
|
||||
*images_out = result_images;
|
||||
} else {
|
||||
free_sd_images(result_images, 1);
|
||||
}
|
||||
return true;
|
||||
}
|
||||
|
||||
int get_upscale_factor(upscaler_ctx_t* upscaler_ctx) {
|
||||
|
||||
Reference in New Issue
Block a user