refactor: return bool from image and upscale APIs (#1728)

This commit is contained in:
leejet 2026-07-01 01:11:56 +08:00 committed by GitHub
parent ccda89e09c
commit 2bb0389683
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
7 changed files with 106 additions and 40 deletions

View File

@ -766,8 +766,12 @@ int main(int argc, const char* argv[]) {
if (cli_params.mode == IMG_GEN) { if (cli_params.mode == IMG_GEN) {
sd_img_gen_params_t img_gen_params = gen_params.to_sd_img_gen_params_t(); sd_img_gen_params_t img_gen_params = gen_params.to_sd_img_gen_params_t();
num_results = gen_params.batch_count; sd_image_t* generated_images = nullptr;
results.adopt(generate_image(sd_ctx.get(), &img_gen_params), num_results); 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) { } else if (cli_params.mode == VID_GEN) {
sd_vid_gen_params_t vid_gen_params = gen_params.to_sd_vid_gen_params_t(); sd_vid_gen_params_t vid_gen_params = gen_params.to_sd_vid_gen_params_t();
sd_image_t* generated_video = nullptr; sd_image_t* generated_video = nullptr;
@ -802,15 +806,21 @@ int main(int argc, const char* argv[]) {
SDImageOwner current_image(results[i]); SDImageOwner current_image(results[i]);
results[i] = {0, 0, 0, nullptr}; results[i] = {0, 0, 0, nullptr};
for (int u = 0; u < gen_params.upscale_repeats; ++u) { for (int u = 0; u < gen_params.upscale_repeats; ++u) {
sd_image_t* upscaled_images = upscale(upscaler_ctx.get(), current_image.get(), upscale_factor); sd_image_t* upscaled_images = nullptr;
if (upscaled_images == nullptr || upscaled_images[0].data == nullptr) { int upscaled_count = 0;
free_sd_images(upscaled_images, 1); 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"); LOG_ERROR("upscale failed");
break; break;
} }
sd_image_t upscaled_image = upscaled_images[0]; sd_image_t upscaled_image = upscaled_images[0];
upscaled_images[0] = {0, 0, 0, nullptr}; 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); current_image.reset(upscaled_image);
} }
results[i] = current_image.release(); // Set the final upscaled image as the result results[i] = current_image.release(); // Set the final upscaled image as the result

View File

@ -173,8 +173,13 @@ bool execute_img_gen_job(ServerRuntime& runtime,
{ {
std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex); std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex);
sd_image_t* raw_results = generate_image(runtime.sd_ctx, &params); sd_image_t* raw_results = nullptr;
results.adopt(raw_results, params.batch_count); int num_results = 0;
if (!generate_image(runtime.sd_ctx, &params, &raw_results, &num_results)) {
raw_results = nullptr;
num_results = 0;
}
results.adopt(raw_results, num_results);
} }
const int num_results = results.count(); const int num_results = results.count();

View File

@ -229,8 +229,11 @@ static bool execute_sync_img_gen_request(ServerRuntime& runtime,
{ {
std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex); std::lock_guard<std::mutex> lock(*runtime.sd_ctx_mutex);
sd_image_t* raw_results = generate_image(runtime.sd_ctx, &img_gen_params); sd_image_t* raw_results = nullptr;
num_results = request.gen_params.batch_count; 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); results.adopt(raw_results, num_results);
} }

View File

@ -292,8 +292,11 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {
{ {
std::lock_guard<std::mutex> lock(*runtime->sd_ctx_mutex); std::lock_guard<std::mutex> lock(*runtime->sd_ctx_mutex);
sd_image_t* raw_results = generate_image(runtime->sd_ctx, &img_gen_params); sd_image_t* raw_results = nullptr;
num_results = request.gen_params.batch_count; 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); results.adopt(raw_results, num_results);
} }

View File

@ -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 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 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 { enum sd_cancel_mode_t {
// Stop the current generation as soon as possible. // 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); const char* params_backend);
SD_API void free_upscaler_ctx(upscaler_ctx_t* upscaler_ctx); SD_API void free_upscaler_ctx(upscaler_ctx_t* upscaler_ctx);
SD_API sd_image_t* upscale(upscaler_ctx_t* upscaler_ctx, SD_API bool upscale(upscaler_ctx_t* upscaler_ctx,
sd_image_t input_image, sd_image_t input_image,
uint32_t upscale_factor); uint32_t upscale_factor,
sd_image_t** images_out,
int* num_images_out);
SD_API int get_upscale_factor(upscaler_ctx_t* upscaler_ctx); SD_API int get_upscale_factor(upscaler_ctx_t* upscaler_ctx);

View File

@ -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, static sd_image_t* decode_image_outputs(sd_ctx_t* sd_ctx,
const GenerationRequest& request, 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()) { if (final_latents.empty()) {
LOG_ERROR("no latent images to decode"); LOG_ERROR("no latent images to decode");
return nullptr; return nullptr;
@ -4320,11 +4321,14 @@ static sd_image_t* decode_image_outputs(sd_ctx_t* sd_ctx,
return nullptr; 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) { if (result_images == nullptr) {
return 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++) { for (size_t i = 0; i < decoded_images.size(); i++) {
result_images[i] = tensor_to_sd_image(decoded_images[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()); 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) { if (sd_ctx == nullptr || sd_img_gen_params == nullptr) {
return nullptr; return false;
} }
sd_ctx->sd->reset_cancel_flag(); 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, &request,
&plan); &plan);
if (!latents_opt.has_value()) { if (!latents_opt.has_value()) {
return nullptr; return false;
} }
ImageGenerationLatents latents = std::move(*latents_opt); 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, &plan,
&latents); &latents);
if (!embeds_opt.has_value()) { if (!embeds_opt.has_value()) {
return nullptr; return false;
} }
ImageGenerationEmbeds embeds = std::move(*embeds_opt); 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(); sd_cancel_mode_t cancel = sd_ctx->sd->get_cancel_flag();
if (cancel == SD_CANCEL_ALL) { if (cancel == SD_CANCEL_ALL) {
LOG_ERROR("cancelling generation"); LOG_ERROR("cancelling generation");
return nullptr; return false;
} }
if (cancel == SD_CANCEL_NEW_LATENTS) { if (cancel == SD_CANCEL_NEW_LATENTS) {
LOG_INFO("cancelling new latent generation, returning %zu/%d completed 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, b + 1,
request.batch_count, request.batch_count,
(sampling_end - sampling_start) * 1.0f / 1000); (sampling_end - sampling_start) * 1.0f / 1000);
return nullptr; return false;
} }
int64_t denoise_end = ggml_time_ms(); int64_t denoise_end = ggml_time_ms();
LOG_INFO("generating %zu latent images completed, taking %.2fs", 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); (denoise_end - denoise_start) * 1.0f / 1000);
if (final_latents.empty()) { if (final_latents.empty()) {
LOG_ERROR("no latent images generated"); LOG_ERROR("no latent images generated");
return nullptr; return false;
} }
if (request.hires.enabled && request.hires.target_width > 0) { if (request.hires.enabled && request.hires.target_width > 0) {
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) { if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
LOG_ERROR("cancelling generation before hires fix"); 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); 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 (request.hires.upscaler == SD_HIRES_UPSCALER_MODEL) {
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) { if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
LOG_ERROR("cancelling generation before hires model load"); 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); LOG_INFO("hires fix: loading model upscaler from '%s'", request.hires.model_path);
hires_upscaler = std::make_unique<UpscalerGGML>(sd_ctx->sd->n_threads, 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, if (!hires_upscaler->load_from_file(request.hires.model_path,
sd_ctx->sd->n_threads)) { sd_ctx->sd->n_threads)) {
LOG_ERROR("load hires model upscaler failed"); 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++) { for (int b = 0; b < (int)final_latents.size(); b++) {
if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) { if (sd_ctx->sd->get_cancel_flag() == SD_CANCEL_ALL) {
LOG_ERROR("cancelling generation during hires fix"); LOG_ERROR("cancelling generation during hires fix");
return nullptr; return false;
} }
int64_t cur_seed = request.seed + b; int64_t cur_seed = request.seed + b;
sd_ctx->sd->rng->manual_seed(cur_seed); 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, request,
hires_upscaler.get()); hires_upscaler.get());
if (upscaled.empty()) { if (upscaled.empty()) {
return nullptr; return false;
} }
sd::Tensor<float> noise = sd::randn_like<float>(upscaled, sd_ctx->sd->rng); 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, b + 1,
(int)final_latents.size(), (int)final_latents.size(),
(hires_sample_end - hires_sample_start) * 1.0f / 1000); (hires_sample_end - hires_sample_start) * 1.0f / 1000);
return nullptr; return false;
} }
int64_t hires_denoise_end = ggml_time_ms(); int64_t hires_denoise_end = ggml_time_ms();
LOG_INFO("hires fix completed, taking %.2fs", (hires_denoise_end - hires_denoise_start) * 1.0f / 1000); 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); 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) { if (result == nullptr) {
return nullptr; return false;
} }
sd_ctx->sd->lora_stat(); sd_ctx->sd->lora_stat();
int64_t t1 = ggml_time_ms(); int64_t t1 = ggml_time_ms();
LOG_INFO("generate_image completed in %.2fs", (t1 - t0) * 1.0f / 1000); 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, static std::optional<ImageGenerationLatents> prepare_video_generation_latents(sd_ctx_t* sd_ctx,

View File

@ -199,23 +199,41 @@ upscaler_ctx_t* new_upscaler_ctx(const char* esrgan_path_c_str,
return upscaler_ctx; 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) { 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)); sd_image_t* result_images = (sd_image_t*)calloc(1, sizeof(sd_image_t));
if (result_images == nullptr) { if (result_images == nullptr) {
return nullptr; return false;
} }
result_images[0] = upscaler_ctx->upscaler->upscale(input_image, upscale_factor); result_images[0] = upscaler_ctx->upscaler->upscale(input_image, upscale_factor);
if (result_images[0].data == nullptr) { if (result_images[0].data == nullptr) {
free(result_images); 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) { int get_upscale_factor(upscaler_ctx_t* upscaler_ctx) {