Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
133 changes: 132 additions & 1 deletion examples/server/routes_sdapi.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -301,6 +301,74 @@ static nlohmann::json prepare_info_field(const SDContextParams& ctx_params,
return jsoninfo;
}

namespace {
struct SdapiPreviewState {
std::mutex mutex;
std::string b64;
std::string id_task;
std::string job_type = "txt2img";
int step = 0;
int total = 0;
bool active = false;
int job_count = 1;
int job_no = 0;
};

SdapiPreviewState g_preview;

}

static void preview_callback(int step, int frame_count, sd_image_t* frames, bool is_noisy, void* data) {
(void)is_noisy;
(void)data;
if (frame_count <= 0 || frames == nullptr || frames[0].data == nullptr) {
return;
}

const sd_image_t& frame = frames[0];

auto image_bytes = encode_image_to_vector(EncodedImageFormat::JPEG,
frame.data,
frame.width,
frame.height,
frame.channel,
"",
80);
if (image_bytes.empty()) {
return;
}

std::lock_guard<std::mutex> lock(g_preview.mutex);
g_preview.b64 = base64_encode(image_bytes);

if (step < g_preview.step) {
if (g_preview.job_no < g_preview.job_count - 1) {
g_preview.job_no++;
}
}
g_preview.step = step;
}

static void set_preview_callback(const std::string& id_task, int steps, int count, const std::string& job_type) {
std::lock_guard<std::mutex> lock(g_preview.mutex);
g_preview.id_task = id_task;
g_preview.total = steps;
g_preview.job_type = job_type;
g_preview.b64.clear();
g_preview.step = 0;
g_preview.active = true;
g_preview.job_count = count > 0 ? count : 1;
g_preview.job_no = 0;
sd_set_preview_callback(preview_callback, PREVIEW_PROJ, 1, true, false, nullptr);
}

static void clear_preview_callback() {
sd_set_preview_callback(nullptr, PREVIEW_NONE, 1, false, false, nullptr);
std::lock_guard<std::mutex> lock(g_preview.mutex);
g_preview.step = g_preview.total;
g_preview.active = false;
}

void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {
ServerRuntime* runtime = &rt;

Expand All @@ -317,7 +385,9 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {
return;
}

json j = json::parse(req.body);
std::string id_task = "sdapi";
json j = json::parse(req.body);
id_task = j.value("id_task", id_task);
ImgGenJobRequest request;
std::string error_message;
if (!build_sdapi_img_gen_request(j, *runtime, img2img, request, error_message)) {
Expand All @@ -334,11 +404,16 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {

{
std::lock_guard<std::mutex> lock(*runtime->sd_ctx_mutex);
set_preview_callback(id_task,
img_gen_params.sample_params.sample_steps,
img_gen_params.batch_count,
img2img ? "img2img" : "txt2img");
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;
}
clear_preview_callback();
results.adopt(raw_results, num_results);
}

Expand Down Expand Up @@ -528,4 +603,60 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) {
r["sd_model_checkpoint"] = model_path.stem();
res.set_content(r.dump(), "application/json");
});

svr.Get("/sdapi/v1/progress", [runtime](const httplib::Request& req, httplib::Response& res) {
bool skip_current_image = false;
bool has_id_param = req.has_param("id_task");
std::string req_id_task = has_id_param ? req.get_param_value("id_task") : "";

if (req.has_param("skip_current_image")) {
std::string skip_str = req.get_param_value("skip_current_image");
skip_current_image = (skip_str == "true" || skip_str == "1");
}

json r;
r["eta_relative"] = 0.0f;
json state;
state["skipped"] = false;
state["interrupted"] = false;
state["job_timestamp"] = "0";
r["textinfo"] = nullptr;

{
std::lock_guard<std::mutex> lock(g_preview.mutex);

bool is_matching_task = !has_id_param || (req_id_task == g_preview.id_task);

int step = is_matching_task ? g_preview.step : 0;
int steps = is_matching_task ? g_preview.total : 0;
std::string b64 = is_matching_task ? g_preview.b64 : "";
std::string job_type = is_matching_task ? g_preview.job_type : "";
int job_count = is_matching_task ? g_preview.job_count : 1;
int job_no = is_matching_task ? g_preview.job_no : 0;

float batch_progress = (steps > 0) ? (static_cast<float>(step) / static_cast<float>(steps)) : 0.0f;
if (batch_progress > 1.0f)
batch_progress = 1.0f;

float overall_progress = (static_cast<float>(job_no) + batch_progress) / static_cast<float>(job_count);
if (overall_progress > 1.0f)
overall_progress = 1.0f;

r["progress"] = overall_progress;
state["job"] = job_type;
state["sampling_step"] = step;
state["sampling_steps"] = steps;
state["job_count"] = job_count;
state["job_no"] = job_no;
r["state"] = state;

if (!skip_current_image && !b64.empty()) {
r["current_image"] = b64;
} else {
r["current_image"] = nullptr;
}
}

res.set_content(r.dump(), "application/json");
});
}
Loading