feat: add sdcpp api support (#1407)

This commit is contained in:
leejet
2026-04-11 17:49:00 +08:00
committed by GitHub
parent 118489eb5c
commit 7ade90e478
17 changed files with 3237 additions and 1308 deletions

View File

@@ -192,17 +192,22 @@ struct SDCliParams {
return options;
};
bool process_and_check() {
if (mode != METADATA && output_path.length() == 0) {
LOG_ERROR("error: the following arguments are required: output_path");
return false;
}
bool resolve() {
if (mode == CONVERT) {
if (output_path == "output.png") {
output_path = "output.gguf";
}
} else if (mode == METADATA) {
}
return true;
}
bool validate() {
if (mode != METADATA) {
if (output_path.length() == 0) {
LOG_ERROR("error: the following arguments are required: output_path");
return false;
}
} else {
if (image_path.empty()) {
LOG_ERROR("error: metadata mode needs an image path (--image)");
return false;
@@ -216,6 +221,16 @@ struct SDCliParams {
return true;
}
bool resolve_and_validate() {
if (!resolve()) {
return false;
}
if (!validate()) {
return false;
}
return true;
}
std::string to_string() const {
std::ostringstream oss;
oss << "SDCliParams {\n"
@@ -260,10 +275,10 @@ void parse_args(int argc, const char** argv, SDCliParams& cli_params, SDContextP
exit(cli_params.normal_exit ? 0 : 1);
}
bool valid = cli_params.process_and_check();
bool valid = cli_params.resolve_and_validate();
if (valid && cli_params.mode != METADATA) {
valid = ctx_params.process_and_check(cli_params.mode) &&
gen_params.process_and_check(cli_params.mode, ctx_params.lora_model_dir);
valid = ctx_params.resolve_and_validate(cli_params.mode) &&
gen_params.resolve_and_validate(cli_params.mode, ctx_params.lora_model_dir);
}
if (!valid) {
@@ -278,7 +293,7 @@ void sd_log_cb(enum sd_log_level_t level, const char* log, void* data) {
}
bool load_images_from_dir(const std::string dir,
SDImageVec& images,
std::vector<SDImageOwner>& images,
int expected_width = 0,
int expected_height = 0,
int max_image_num = 0,
@@ -315,10 +330,10 @@ bool load_images_from_dir(const std::string dir,
return false;
}
images.push_back({(uint32_t)width,
(uint32_t)height,
3,
image_buffer});
images.emplace_back(sd_image_t{(uint32_t)width,
(uint32_t)height,
3,
image_buffer});
if (max_image_num > 0 && static_cast<int>(images.size()) >= max_image_num) {
break;
@@ -558,13 +573,6 @@ int main(int argc, const char* argv[]) {
}
bool vae_decode_only = true;
SDImageOwner init_image({0, 0, 3, nullptr});
SDImageOwner end_image({0, 0, 3, nullptr});
SDImageOwner control_image({0, 0, 3, nullptr});
SDImageOwner mask_image({0, 0, 1, nullptr});
SDImageVec ref_images;
SDImageVec pmid_images;
SDImageVec control_frames;
auto load_image_and_update_size = [&](const std::string& path,
SDImageOwner& image,
@@ -588,31 +596,32 @@ int main(int argc, const char* argv[]) {
if (gen_params.init_image_path.size() > 0) {
vae_decode_only = false;
if (!load_image_and_update_size(gen_params.init_image_path, init_image)) {
if (!load_image_and_update_size(gen_params.init_image_path, gen_params.init_image)) {
return 1;
}
}
if (gen_params.end_image_path.size() > 0) {
vae_decode_only = false;
if (!load_image_and_update_size(gen_params.end_image_path, end_image)) {
if (!load_image_and_update_size(gen_params.end_image_path, gen_params.end_image)) {
return 1;
}
}
if (gen_params.ref_image_paths.size() > 0) {
vae_decode_only = false;
gen_params.ref_images.clear();
for (auto& path : gen_params.ref_image_paths) {
SDImageOwner ref_image({0, 0, 3, nullptr});
if (!load_image_and_update_size(path, ref_image, false)) {
return 1;
}
ref_images.push_back(std::move(ref_image));
gen_params.ref_images.push_back(std::move(ref_image));
}
}
if (gen_params.mask_image_path.size() > 0) {
if (!load_sd_image_from_file(mask_image.put(),
if (!load_sd_image_from_file(gen_params.mask_image.put(),
gen_params.mask_image_path.c_str(),
gen_params.get_resolved_width(),
gen_params.get_resolved_height(),
@@ -630,11 +639,11 @@ int main(int argc, const char* argv[]) {
generated_mask.width = gen_params.get_resolved_width();
generated_mask.height = gen_params.get_resolved_height();
memset(generated_mask.data, 255, gen_params.get_resolved_width() * gen_params.get_resolved_height());
mask_image.reset(generated_mask);
gen_params.mask_image.reset(generated_mask);
}
if (gen_params.control_image_path.size() > 0) {
if (!load_sd_image_from_file(control_image.put(),
if (!load_sd_image_from_file(gen_params.control_image.put(),
gen_params.control_image_path.c_str(),
gen_params.get_resolved_width(),
gen_params.get_resolved_height())) {
@@ -642,7 +651,7 @@ int main(int argc, const char* argv[]) {
return 1;
}
if (cli_params.canny_preprocess) { // apply preprocessor
preprocess_canny(control_image.get(),
preprocess_canny(gen_params.control_image.get(),
0.08f,
0.08f,
0.8f,
@@ -652,8 +661,9 @@ int main(int argc, const char* argv[]) {
}
if (!gen_params.control_video_path.empty()) {
gen_params.control_frames.clear();
if (!load_images_from_dir(gen_params.control_video_path,
control_frames,
gen_params.control_frames,
gen_params.get_resolved_width(),
gen_params.get_resolved_height(),
gen_params.video_frames,
@@ -663,8 +673,9 @@ int main(int argc, const char* argv[]) {
}
if (!gen_params.pm_id_images_dir.empty()) {
gen_params.pm_id_images.clear();
if (!load_images_from_dir(gen_params.pm_id_images_dir,
pmid_images,
gen_params.pm_id_images,
0,
0,
0,
@@ -684,7 +695,7 @@ int main(int argc, const char* argv[]) {
if (cli_params.mode == UPSCALE) {
num_results = 1;
results.push_back(init_image.release());
results.push_back(gen_params.init_image.release());
} else {
SDCtxPtr sd_ctx(new_sd_ctx(&sd_ctx_params));
@@ -706,63 +717,13 @@ int main(int argc, const char* argv[]) {
}
if (cli_params.mode == IMG_GEN) {
sd_img_gen_params_t img_gen_params = {
gen_params.lora_vec.data(),
static_cast<uint32_t>(gen_params.lora_vec.size()),
gen_params.prompt.c_str(),
gen_params.negative_prompt.c_str(),
gen_params.clip_skip,
init_image.get(),
ref_images.data(),
(int)ref_images.size(),
gen_params.auto_resize_ref_image,
gen_params.increase_ref_index,
mask_image.get(),
gen_params.get_resolved_width(),
gen_params.get_resolved_height(),
gen_params.sample_params,
gen_params.strength,
gen_params.seed,
gen_params.batch_count,
control_image.get(),
gen_params.control_strength,
{
pmid_images.data(),
(int)pmid_images.size(),
gen_params.pm_id_embed_path.c_str(),
gen_params.pm_style_strength,
}, // pm_params
gen_params.vae_tiling_params,
gen_params.cache_params,
};
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);
} else if (cli_params.mode == VID_GEN) {
sd_vid_gen_params_t vid_gen_params = {
gen_params.lora_vec.data(),
static_cast<uint32_t>(gen_params.lora_vec.size()),
gen_params.prompt.c_str(),
gen_params.negative_prompt.c_str(),
gen_params.clip_skip,
init_image.get(),
end_image.get(),
control_frames.data(),
(int)control_frames.size(),
gen_params.get_resolved_width(),
gen_params.get_resolved_height(),
gen_params.sample_params,
gen_params.high_noise_sample_params,
gen_params.moe_boundary,
gen_params.strength,
gen_params.seed,
gen_params.video_frames,
gen_params.vace_strength,
gen_params.vae_tiling_params,
gen_params.cache_params,
};
sd_image_t* generated_video = generate_video(sd_ctx.get(), &vid_gen_params, &num_results);
sd_vid_gen_params_t vid_gen_params = gen_params.to_sd_vid_gen_params_t();
sd_image_t* generated_video = generate_video(sd_ctx.get(), &vid_gen_params, &num_results);
results.adopt(generated_video, num_results);
}