ACMX 2.139.0
Dual-Backend Real-Time GPU Video Synthesis
Loading...
Searching...
No Matches
stable_diffusion.cpp
Go to the documentation of this file.
2
3#include <curl/curl.h>
4#include <json/json.h>
5#include <opencv2/imgcodecs.hpp>
6#include <opencv2/imgproc.hpp>
7
8#include <algorithm>
9#include <chrono>
10#include <cmath>
11#ifdef _WIN32
12#ifndef NOMINMAX
13#define NOMINMAX
14#endif
15#include <windows.h>
16#else
17#include <cerrno>
18#include <csignal>
19#include <cstring>
20#include <fcntl.h>
21#endif
22#include <fstream>
23#include <iostream>
24#include <limits>
25#include <stdexcept>
26#include <string_view>
27#include <thread>
28#include <utility>
29#include <vector>
30
31#ifndef _WIN32
32#include <spawn.h>
33#include <sys/types.h>
34#include <sys/wait.h>
35#include <unistd.h>
36extern char **environ;
37#endif
38
40 namespace {
41 constexpr std::size_t MAX_RESPONSE_SIZE = 64U * 1024U * 1024U;
42
44 std::string value;
45 bool overflow = false;
46 };
47
48 [[nodiscard]] std::string curlError(CURLcode code) { return curl_easy_strerror(code); }
49
50#ifdef _WIN32
51 [[nodiscard]] std::string windowsError(DWORD error) {
52 LPWSTR message = nullptr;
53 const DWORD length = FormatMessageW(FORMAT_MESSAGE_ALLOCATE_BUFFER | FORMAT_MESSAGE_FROM_SYSTEM | FORMAT_MESSAGE_IGNORE_INSERTS, nullptr, error, 0, reinterpret_cast<LPWSTR>(&message), 0, nullptr);
54 std::string text;
55 if (length != 0U && message != nullptr) {
56 const int utf8_length = WideCharToMultiByte(CP_UTF8, 0, message, static_cast<int>(length), nullptr, 0, nullptr, nullptr);
57 if (utf8_length > 0) {
58 text.resize(static_cast<std::size_t>(utf8_length));
59 WideCharToMultiByte(CP_UTF8, 0, message, static_cast<int>(length), text.data(), utf8_length, nullptr, nullptr);
60 }
61 }
62 if (message != nullptr) {
63 LocalFree(message);
64 }
65 while (!text.empty() && (text.back() == '\r' || text.back() == '\n')) {
66 text.pop_back();
67 }
68 return text.empty() ? "Windows error " + std::to_string(error) : text;
69 }
70
71 [[nodiscard]] std::wstring utf8ToWide(std::string_view text) {
72 if (text.empty()) {
73 return {};
74 }
75 const int length = MultiByteToWideChar(CP_UTF8, MB_ERR_INVALID_CHARS, text.data(), static_cast<int>(text.size()), nullptr, 0);
76 if (length <= 0) {
77 throw std::runtime_error("unable to convert sd-server argument to UTF-16: " + windowsError(GetLastError()));
78 }
79 std::wstring output(static_cast<std::size_t>(length), L'\0');
80 if (MultiByteToWideChar(CP_UTF8, MB_ERR_INVALID_CHARS, text.data(), static_cast<int>(text.size()), output.data(), length) <= 0) {
81 throw std::runtime_error("unable to convert sd-server argument to UTF-16: " + windowsError(GetLastError()));
82 }
83 return output;
84 }
85
86 [[nodiscard]] std::wstring quoteWindowsArgument(std::wstring_view argument) {
87 if (!argument.empty() && argument.find_first_of(L" \t\n\v\"") == std::wstring_view::npos) {
88 return std::wstring(argument);
89 }
90 std::wstring quoted(L"\"");
91 std::size_t backslashes = 0;
92 for (const wchar_t character : argument) {
93 if (character == L'\\') {
94 ++backslashes;
95 continue;
96 }
97 if (character == L'\"') {
98 quoted.append(backslashes * 2U + 1U, L'\\');
99 quoted.push_back(L'\"');
100 backslashes = 0;
101 continue;
102 }
103 quoted.append(backslashes, L'\\');
104 backslashes = 0;
105 quoted.push_back(character);
106 }
107 quoted.append(backslashes * 2U, L'\\');
108 quoted.push_back(L'\"');
109 return quoted;
110 }
111#endif
112
113 std::size_t appendResponse(char *data, std::size_t size, std::size_t count, void *context) {
114 auto *buffer = static_cast<ResponseBuffer *>(context);
115 if (size != 0U && count > std::numeric_limits<std::size_t>::max() / size) {
116 buffer->overflow = true;
117 return 0U;
118 }
119 const std::size_t bytes = size * count;
120 if (bytes > MAX_RESPONSE_SIZE - std::min(buffer->value.size(), MAX_RESPONSE_SIZE)) {
121 buffer->overflow = true;
122 return 0U;
123 }
124 buffer->value.append(data, bytes);
125 return bytes;
126 }
127
128 int checkCancelled(void *context, curl_off_t, curl_off_t, curl_off_t, curl_off_t) {
129 const auto *cancelled = static_cast<const std::function<bool()> *>(context);
130 return cancelled != nullptr && *cancelled && (*cancelled)() ? 1 : 0;
131 }
132
133 [[nodiscard]] std::string encodeBase64(const std::vector<std::uint8_t> &bytes) {
134 static constexpr std::string_view ALPHABET = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
135 std::string output;
136 output.reserve(((bytes.size() + 2U) / 3U) * 4U);
137 for (std::size_t index = 0; index < bytes.size(); index += 3U) {
138 const std::uint32_t first = bytes[index];
139 const std::uint32_t second = index + 1U < bytes.size() ? bytes[index + 1U] : 0U;
140 const std::uint32_t third = index + 2U < bytes.size() ? bytes[index + 2U] : 0U;
141 const std::uint32_t value = (first << 16U) | (second << 8U) | third;
142 output.push_back(ALPHABET[(value >> 18U) & 0x3fU]);
143 output.push_back(ALPHABET[(value >> 12U) & 0x3fU]);
144 output.push_back(index + 1U < bytes.size() ? ALPHABET[(value >> 6U) & 0x3fU] : '=');
145 output.push_back(index + 2U < bytes.size() ? ALPHABET[value & 0x3fU] : '=');
146 }
147 return output;
148 }
149
150 [[nodiscard]] int decodeBase64Character(unsigned char character) {
151 if (character >= 'A' && character <= 'Z') {
152 return character - 'A';
153 }
154 if (character >= 'a' && character <= 'z') {
155 return character - 'a' + 26;
156 }
157 if (character >= '0' && character <= '9') {
158 return character - '0' + 52;
159 }
160 if (character == '+') {
161 return 62;
162 }
163 if (character == '/') {
164 return 63;
165 }
166 return -1;
167 }
168
169 [[nodiscard]] std::vector<std::uint8_t> decodeBase64(std::string_view encoded) {
170 const std::size_t comma = encoded.find(',');
171 if (encoded.starts_with("data:") && comma != std::string_view::npos) {
172 encoded.remove_prefix(comma + 1U);
173 }
174 std::vector<std::uint8_t> output;
175 output.reserve((encoded.size() / 4U) * 3U);
176 std::uint32_t accumulator = 0U;
177 int bits = 0;
178 for (const unsigned char character : encoded) {
179 if (character == '=') {
180 break;
181 }
182 const int value = decodeBase64Character(character);
183 if (value < 0) {
184 if (character == ' ' || character == '\n' || character == '\r' || character == '\t') {
185 continue;
186 }
187 throw std::runtime_error("sd-server returned invalid base64 image data");
188 }
189 accumulator = (accumulator << 6U) | static_cast<std::uint32_t>(value);
190 bits += 6;
191 if (bits >= 8) {
192 bits -= 8;
193 output.push_back(static_cast<std::uint8_t>((accumulator >> static_cast<unsigned int>(bits)) & 0xffU));
194 }
195 }
196 return output;
197 }
198
199 [[nodiscard]] Json::Value parseJson(std::string_view text, std::string_view context) {
200 Json::CharReaderBuilder builder;
201 builder["collectComments"] = false;
202 std::unique_ptr<Json::CharReader> reader(builder.newCharReader());
203 Json::Value value;
204 std::string error;
205 if (!reader->parse(text.data(), text.data() + text.size(), &value, &error)) {
206 throw std::runtime_error(std::string(context) + " returned invalid JSON: " + error);
207 }
208 return value;
209 }
210
211 [[nodiscard]] std::string writeJson(const Json::Value &value) {
212 Json::StreamWriterBuilder builder;
213 builder["indentation"] = "";
214 return Json::writeString(builder, value);
215 }
216
217 [[nodiscard]] ResponseBuffer request(std::string_view url, const std::string *body, long timeout_seconds, long &status, const std::function<bool()> *cancelled) {
218 CURL *handle = curl_easy_init();
219 if (handle == nullptr) {
220 throw std::runtime_error("unable to initialize libcurl");
221 }
222 ResponseBuffer response;
223 curl_slist *headers = nullptr;
224 if (body != nullptr) {
225 headers = curl_slist_append(headers, "Content-Type: application/json");
226 }
227 const std::string request_url(url);
228 curl_easy_setopt(handle, CURLOPT_URL, request_url.c_str());
229 curl_easy_setopt(handle, CURLOPT_NOPROXY, "127.0.0.1,localhost");
230 curl_easy_setopt(handle, CURLOPT_WRITEFUNCTION, appendResponse);
231 curl_easy_setopt(handle, CURLOPT_WRITEDATA, &response);
232 curl_easy_setopt(handle, CURLOPT_CONNECTTIMEOUT, 2L);
233 curl_easy_setopt(handle, CURLOPT_TIMEOUT, timeout_seconds);
234 curl_easy_setopt(handle, CURLOPT_NOSIGNAL, 1L);
235 curl_easy_setopt(handle, CURLOPT_NOPROGRESS, 0L);
236 curl_easy_setopt(handle, CURLOPT_XFERINFOFUNCTION, checkCancelled);
237 curl_easy_setopt(handle, CURLOPT_XFERINFODATA, cancelled);
238 if (body != nullptr) {
239 curl_easy_setopt(handle, CURLOPT_HTTPHEADER, headers);
240 curl_easy_setopt(handle, CURLOPT_POST, 1L);
241 curl_easy_setopt(handle, CURLOPT_POSTFIELDS, body->data());
242 curl_easy_setopt(handle, CURLOPT_POSTFIELDSIZE_LARGE, static_cast<curl_off_t>(body->size()));
243 }
244 const CURLcode result = curl_easy_perform(handle);
245 curl_easy_getinfo(handle, CURLINFO_RESPONSE_CODE, &status);
246 curl_slist_free_all(headers);
247 curl_easy_cleanup(handle);
248 if (result != CURLE_OK) {
249 if (result == CURLE_ABORTED_BY_CALLBACK) {
250 throw std::runtime_error("Stable Diffusion request cancelled");
251 }
252 if (response.overflow) {
253 throw std::runtime_error("sd-server response exceeded 64 MiB");
254 }
255 throw std::runtime_error("sd-server request failed: " + curlError(result));
256 }
257 return response;
258 }
259
260 [[nodiscard]] std::string serverError(const Json::Value &root) {
261 if (root.isMember("message") && root["message"].isString()) {
262 return root["message"].asString();
263 }
264 if (root.isMember("error") && root["error"].isString()) {
265 return root["error"].asString();
266 }
267 return "unknown server error";
268 }
269
270 [[nodiscard]] cv::Size neuralUpscaleWorkingSize(const Settings &settings, const cv::Size &fallback) {
271 if (settings.upscale_working_width > 0 && settings.upscale_working_height > 0) {
272 return {settings.upscale_working_width, settings.upscale_working_height};
273 }
274 int width = settings.upscale_width > 0 ? settings.upscale_width : fallback.width;
275 int height = settings.upscale_height > 0 ? settings.upscale_height : fallback.height;
276 constexpr double MAX_WORKING_PIXELS = 1280.0 * 720.0;
277 const double pixels = static_cast<double>(width) * height;
278 if (pixels > MAX_WORKING_PIXELS) {
279 const double scale = std::sqrt(MAX_WORKING_PIXELS / pixels);
280 width = std::max(64, static_cast<int>(std::floor(width * scale)));
281 height = std::max(64, static_cast<int>(std::floor(height * scale)));
282 }
283 return {width, height};
284 }
285
286 void appendLoras(Json::Value &root, const Settings &settings) {
287 if (settings.loras.empty()) {
288 return;
289 }
290 root["lora"] = Json::arrayValue;
291 for (const Lora &lora : settings.loras) {
292 Json::Value entry;
293 entry["path"] = lora.path.generic_string();
294 entry["multiplier"] = lora.multiplier;
295 entry["is_high_noise"] = false;
296 root["lora"].append(std::move(entry));
297 }
298 }
299
300 [[nodiscard]] std::string submitAsyncJob(const std::string &endpoint, std::string_view route, const Json::Value &request_root, const Settings &settings) {
301 const std::string body = writeJson(request_root);
302 long status = 0;
303 const ResponseBuffer submission = request(endpoint + std::string(route), &body, 30L, status, &settings.cancelled);
304 const Json::Value accepted = parseJson(submission.value, "sd-server");
305 if (status != 202 || !accepted["id"].isString()) {
306 throw std::runtime_error("sd-server rejected frame: " + serverError(accepted));
307 }
308 const std::string job_id = accepted["id"].asString();
309 const std::string job_url = endpoint + "/sdcpp/v1/jobs/" + job_id;
310 const auto deadline = std::chrono::steady_clock::now() + std::chrono::hours(1);
311 while (std::chrono::steady_clock::now() < deadline) {
312 if (settings.cancelled && settings.cancelled()) {
313 const std::string cancel_body = "{}";
314 long cancel_status = 0;
315 try {
316 static_cast<void>(request(job_url + "/cancel", &cancel_body, 3L, cancel_status, nullptr));
317 } catch (const std::exception &) {
318 }
319 throw std::runtime_error("Stable Diffusion request cancelled");
320 }
321 long poll_status = 0;
322 const ResponseBuffer poll = request(job_url, nullptr, 10L, poll_status, &settings.cancelled);
323 const Json::Value job = parseJson(poll.value, "sd-server job");
324 if (poll_status != 200) {
325 throw std::runtime_error("sd-server job polling failed: " + serverError(job));
326 }
327 const std::string state = job["status"].asString();
328 if (state == "completed") {
329 const Json::Value &images = job["result"]["images"];
330 if (!images.isArray() || images.empty() || !images[0]["b64_json"].isString()) {
331 throw std::runtime_error("sd-server job did not contain an output image");
332 }
333 return images[0]["b64_json"].asString();
334 }
335 if (state == "failed" || state == "cancelled") {
336 throw std::runtime_error("sd-server frame job " + state + ": " + serverError(job["error"]));
337 }
338 std::this_thread::sleep_for(std::chrono::milliseconds(50));
339 }
340 throw std::runtime_error("sd-server frame job timed out");
341 }
342 } // namespace
343
345 if (curl_global_init(CURL_GLOBAL_DEFAULT) != CURLE_OK) {
346 throw std::runtime_error("unable to initialize libcurl runtime");
347 }
348 endpoint = "http://127.0.0.1:" + std::to_string(this->settings.port);
349 try {
350 start();
352 } catch (...) {
353 stop();
354 curl_global_cleanup();
355 throw;
356 }
357 }
358
360 stop();
361 curl_global_cleanup();
362 }
363
365 const std::string port = std::to_string(settings.port);
366 const std::string executable = settings.server_executable.string();
367 std::vector<std::string> argument_storage{executable, "--listen-ip", "127.0.0.1", "--listen-port", port};
369 argument_storage.emplace_back("--upscale-model");
370 argument_storage.push_back(settings.upscale_model.string());
371 } else {
372 argument_storage.insert(argument_storage.end(), {"--model", settings.model.string(), "--type", "f16", "--mmap", "--fa", "--diffusion-conv-direct", "--vae-conv-direct", "--lora-model-dir", settings.lora_directory.string()});
373 }
374 if (!settings.upscale_only && !settings.upscale_model.empty()) {
375 argument_storage.emplace_back("--hires-upscalers-dir");
376 argument_storage.push_back(settings.upscale_model.parent_path().string());
377 }
378 argument_storage.insert(argument_storage.end(), settings.server_arguments.begin(), settings.server_arguments.end());
379#ifdef _WIN32
380 std::wstring command_line;
381 for (const std::string &argument : argument_storage) {
382 if (!command_line.empty()) {
383 command_line.push_back(L' ');
384 }
385 command_line += quoteWindowsArgument(utf8ToWide(argument));
386 }
387
388 STARTUPINFOW startup_info{};
389 startup_info.cb = sizeof(startup_info);
390 PROCESS_INFORMATION process_info{};
391 HANDLE log_handle = INVALID_HANDLE_VALUE;
392 HANDLE input_handle = INVALID_HANDLE_VALUE;
393 if (settings.quiet) {
394 std::vector<wchar_t> log_path(MAX_PATH, L'\0');
395 if (GetTempFileNameW(std::filesystem::temp_directory_path().wstring().c_str(), L"acm", 0, log_path.data()) == 0) {
396 throw std::runtime_error("unable to create sd-server diagnostic log: " + windowsError(GetLastError()));
397 }
398 diagnostic_log_path = std::filesystem::path(log_path.data());
399 log_handle = CreateFileW(diagnostic_log_path.wstring().c_str(), FILE_APPEND_DATA, FILE_SHARE_READ | FILE_SHARE_WRITE, nullptr, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL, nullptr);
400 if (log_handle == INVALID_HANDLE_VALUE) {
401 throw std::runtime_error("unable to open sd-server diagnostic log: " + windowsError(GetLastError()));
402 }
403 input_handle = CreateFileW(L"NUL", GENERIC_READ, FILE_SHARE_READ | FILE_SHARE_WRITE, nullptr, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL, nullptr);
404 if (input_handle == INVALID_HANDLE_VALUE) {
405 const DWORD error = GetLastError();
406 CloseHandle(log_handle);
407 throw std::runtime_error("unable to configure sd-server input: " + windowsError(error));
408 }
409 if (!SetHandleInformation(log_handle, HANDLE_FLAG_INHERIT, HANDLE_FLAG_INHERIT) || !SetHandleInformation(input_handle, HANDLE_FLAG_INHERIT, HANDLE_FLAG_INHERIT)) {
410 const DWORD error = GetLastError();
411 CloseHandle(input_handle);
412 CloseHandle(log_handle);
413 throw std::runtime_error("unable to configure sd-server output: " + windowsError(error));
414 }
415 startup_info.dwFlags = STARTF_USESTDHANDLES;
416 startup_info.hStdOutput = log_handle;
417 startup_info.hStdError = log_handle;
418 startup_info.hStdInput = input_handle;
419 }
420 std::vector<wchar_t> mutable_command(command_line.begin(), command_line.end());
421 mutable_command.push_back(L'\0');
422 const BOOL launched = CreateProcessW(nullptr, mutable_command.data(), nullptr, nullptr, settings.quiet ? TRUE : FALSE, CREATE_NO_WINDOW, nullptr, nullptr, &startup_info, &process_info);
423 const DWORD launch_error = launched ? ERROR_SUCCESS : GetLastError();
424 if (log_handle != INVALID_HANDLE_VALUE) {
425 CloseHandle(log_handle);
426 }
427 if (input_handle != INVALID_HANDLE_VALUE) {
428 CloseHandle(input_handle);
429 }
430 if (!launched) {
431 throw std::runtime_error("unable to launch sd-server: " + windowsError(launch_error));
432 }
433 CloseHandle(process_info.hThread);
434 process_handle = reinterpret_cast<std::intptr_t>(process_info.hProcess);
435 process_id = static_cast<std::int64_t>(process_info.dwProcessId);
436#else
437 std::vector<char *> arguments;
438 arguments.reserve(argument_storage.size() + 1U);
439 for (std::string &argument : argument_storage) {
440 arguments.push_back(argument.data());
441 }
442 arguments.push_back(nullptr);
443 pid_t child = -1;
444 posix_spawn_file_actions_t file_actions;
445 posix_spawn_file_actions_t *actions = nullptr;
446 if (settings.quiet) {
447 std::string log_template = (std::filesystem::temp_directory_path() / "acmxvk-sd-server-XXXXXX").string();
448 std::vector<char> mutable_template(log_template.begin(), log_template.end());
449 mutable_template.push_back('\0');
450 const int log_fd = ::mkstemp(mutable_template.data());
451 if (log_fd < 0) {
452 throw std::runtime_error("unable to create sd-server diagnostic log: " + std::string(std::strerror(errno)));
453 }
454 ::close(log_fd);
455 diagnostic_log_path = mutable_template.data();
456
457 const int init_result = posix_spawn_file_actions_init(&file_actions);
458 if (init_result != 0) {
459 throw std::runtime_error("unable to configure sd-server output: " + std::string(std::strerror(init_result)));
460 }
461 actions = &file_actions;
462 const int stdout_result = posix_spawn_file_actions_addopen(actions, STDOUT_FILENO, diagnostic_log_path.c_str(), O_WRONLY | O_APPEND, 0600);
463 const int stderr_result = posix_spawn_file_actions_adddup2(actions, STDOUT_FILENO, STDERR_FILENO);
464 if (stdout_result != 0 || stderr_result != 0) {
465 posix_spawn_file_actions_destroy(actions);
466 const int error = stdout_result != 0 ? stdout_result : stderr_result;
467 throw std::runtime_error("unable to redirect sd-server output: " + std::string(std::strerror(error)));
468 }
469 }
470 const int result = posix_spawnp(&child, executable.c_str(), actions, nullptr, arguments.data(), environ);
471 if (actions != nullptr) {
472 posix_spawn_file_actions_destroy(actions);
473 }
474 if (result != 0) {
475 throw std::runtime_error("unable to launch sd-server: " + std::string(std::strerror(result)));
476 }
477 process_id = child;
478#endif
479 std::cout << "acmxvk: launched sd-server process " << process_id << " on 127.0.0.1:" << settings.port << '\n';
480 if (!diagnostic_log_path.empty()) {
481 std::cout << "acmxvk: sd-server diagnostic log: " << diagnostic_log_path << '\n';
482 }
483 }
484
485 void Server::resetDiagnosticLog() const noexcept {
486 if (diagnostic_log_path.empty()) {
487 return;
488 }
489 std::error_code error;
490 std::filesystem::resize_file(diagnostic_log_path, 0U, error);
491 }
492
493 std::string Server::diagnosticLogDetails() const {
494 if (diagnostic_log_path.empty()) {
495 return {};
496 }
497 constexpr std::streamoff MAX_LOG_TAIL = 16 * 1024;
498 std::ifstream input(diagnostic_log_path, std::ios::binary | std::ios::ate);
499 std::string details = "; sd-server diagnostic log: " + diagnostic_log_path.string();
500 if (!input) {
501 return details;
502 }
503 const std::streamoff size = input.tellg();
504 if (size <= 0) {
505 return details;
506 }
507 const std::streamoff start = std::max<std::streamoff>(0, size - MAX_LOG_TAIL);
508 input.seekg(start);
509 std::string tail(static_cast<std::size_t>(size - start), '\0');
510 input.read(tail.data(), static_cast<std::streamsize>(tail.size()));
511 return details + "\n--- sd-server log tail ---\n" + tail;
512 }
513
514 void Server::stop() noexcept {
515#ifdef _WIN32
516 const HANDLE process = reinterpret_cast<HANDLE>(process_handle);
517 if (process == nullptr) {
518 return;
519 }
520 if (WaitForSingleObject(process, 1000U) == WAIT_TIMEOUT && !TerminateProcess(process, 0U)) {
521 std::cerr << "acmxvk: unable to stop sd-server process " << process_id << ": " << windowsError(GetLastError()) << '\n';
522 }
523 WaitForSingleObject(process, 1000U);
524 CloseHandle(process);
525 process_handle = 0;
526 process_id = -1;
527#else
528 if (process_id <= 0) {
529 return;
530 }
531 const pid_t child = static_cast<pid_t>(process_id);
532 if (kill(child, SIGTERM) != 0 && errno != ESRCH) {
533 std::cerr << "acmxvk: unable to stop sd-server process " << process_id << ": " << std::strerror(errno) << '\n';
534 }
535 for (int attempt = 0; attempt < 50; ++attempt) {
536 int status = 0;
537 const pid_t result = waitpid(child, &status, WNOHANG);
538 if (result == child || (result < 0 && errno == ECHILD)) {
539 process_id = -1;
540 return;
541 }
542 std::this_thread::sleep_for(std::chrono::milliseconds(20));
543 }
544 if (kill(child, SIGKILL) == 0 || errno == ESRCH) {
545 int status = 0;
546 while (waitpid(child, &status, 0) < 0 && errno == EINTR) {
547 }
548 }
549 process_id = -1;
550#endif
551 }
552
554 const bool require_capabilities = settings.upscale_only || !settings.upscale_model.empty() || !settings.loras.empty();
555 const std::string url = endpoint + (require_capabilities ? "/sdcpp/v1/capabilities" : "/sdapi/v1/options");
556 const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(5);
557 while (std::chrono::steady_clock::now() < deadline) {
559 throw std::runtime_error("Stable Diffusion startup cancelled");
560 }
561 bool server_exited = false;
562#ifdef _WIN32
563 const HANDLE process = reinterpret_cast<HANDLE>(process_handle);
564 const DWORD wait_result = process == nullptr ? WAIT_OBJECT_0 : WaitForSingleObject(process, 0U);
565 if (wait_result == WAIT_FAILED) {
566 throw std::runtime_error("unable to inspect sd-server process: " + windowsError(GetLastError()));
567 }
568 server_exited = wait_result == WAIT_OBJECT_0;
569 if (server_exited && process != nullptr) {
570 CloseHandle(process);
571 process_handle = 0;
572 }
573#else
574 int child_status = 0;
575 const pid_t result = waitpid(static_cast<pid_t>(process_id), &child_status, WNOHANG);
576 server_exited = result == static_cast<pid_t>(process_id);
577#endif
578 if (server_exited) {
579 process_id = -1;
580 throw std::runtime_error("sd-server exited before accepting requests" + diagnosticLogDetails());
581 }
582 long status = 0;
583 try {
584 const ResponseBuffer response = request(url, nullptr, 3L, status, &settings.cancelled);
585 if (status == 200) {
586 if (!response.value.empty()) {
587 const Json::Value document = parseJson(response.value, "sd-server");
588 if (!settings.upscale_model.empty()) {
589 const std::string upscaler_name = settings.upscale_model.stem().string();
590 bool found = false;
591 for (const Json::Value &entry : document["upscalers"]) {
592 if (entry["name"].asString() == upscaler_name) {
593 found = true;
594 break;
595 }
596 }
597 if (!found) {
598 throw std::runtime_error("sd-server did not discover the requested "
599 "upscaler model: " +
600 upscaler_name);
601 }
602 }
603 for (const Lora &lora : settings.loras) {
604 bool found = false;
605 for (const Json::Value &entry : document["loras"]) {
606 if (entry["path"].asString() == lora.path.generic_string()) {
607 found = true;
608 break;
609 }
610 }
611 if (!found) {
612 throw std::runtime_error("sd-server did not discover the requested "
613 "LoRA model: " +
614 lora.path.generic_string());
615 }
616 }
617 }
618 std::cout << "acmxvk: sd-server " << (settings.upscale_only ? "upscaler" : "model") << " ready";
619 if (!settings.upscale_only) {
620 std::cout << "; processing " << settings.width << 'x' << settings.height << " frames at " << settings.steps << " configured steps, strength " << settings.strength;
621 }
622 if (!settings.upscale_model.empty()) {
623 const cv::Size working = neuralUpscaleWorkingSize(settings, {settings.width, settings.height});
624 std::cout << "; ESRGAN working resolution " << working.width << 'x' << working.height << ", final resolution " << settings.upscale_width << 'x' << settings.upscale_height;
625 }
626 if (!settings.loras.empty()) {
627 std::cout << "; " << settings.loras.size() << " LoRA model(s)";
628 }
629 std::cout << '\n';
631 return;
632 }
633 if (status >= 400 && status < 500) {
634 throw std::runtime_error("sd-server readiness check returned HTTP " + std::to_string(status) + diagnosticLogDetails());
635 }
636 } catch (const std::exception &) {
637 if ((status > 0 && status < 500) || (settings.cancelled && settings.cancelled())) {
638 throw;
639 }
640 }
641 std::this_thread::sleep_for(std::chrono::milliseconds(250));
642 }
643 throw std::runtime_error("sd-server did not become ready within five minutes" + diagnosticLogDetails());
644 }
645
646 cv::Mat Server::process(const cv::Mat &rgba) const {
647 if (rgba.empty() || rgba.type() != CV_8UC4) {
648 throw std::runtime_error("Stable Diffusion input must be a non-empty RGBA8 frame");
649 }
650 cv::Mat bgr;
651 cv::cvtColor(rgba, bgr, cv::COLOR_RGBA2BGR);
652 std::vector<std::uint8_t> png;
653 if (!cv::imencode(".png", bgr, png)) {
654 throw std::runtime_error("unable to encode Stable Diffusion input frame");
655 }
656
658
659 const std::string encoded_input = encodeBase64(png);
660 Json::Value root;
661 std::string encoded_output;
663 const cv::Size working_size = neuralUpscaleWorkingSize(settings, rgba.size());
664 cv::Mat working_input = rgba;
665 if (working_input.size() != working_size) {
666 cv::resize(working_input, working_input, working_size, 0.0, 0.0, cv::INTER_LANCZOS4);
667 }
668 cv::Mat working_bgr;
669 cv::cvtColor(working_input, working_bgr, cv::COLOR_RGBA2BGR);
670 std::vector<std::uint8_t> working_png;
671 if (!cv::imencode(".png", working_bgr, working_png)) {
672 throw std::runtime_error("unable to encode Stable Diffusion ESRGAN input frame");
673 }
674 root["image"] = encodeBase64(working_png);
675 root["upscaler"] = settings.upscale_model.stem().string();
676 root["upscale_repeats"] = 1;
677 root["upscale_tile_size"] = 128;
678 root["output_compression"] = 100;
679 encoded_output = submitAsyncJob(endpoint, "/sdcpp/v1/upscale", root, settings);
680 } else if (settings.upscale_model.empty()) {
681 root["prompt"] = settings.prompt;
682 root["negative_prompt"] = settings.negative_prompt;
683 root["width"] = settings.width;
684 root["height"] = settings.height;
685 root["seed"] = settings.seed;
686 appendLoras(root, settings);
687 root["steps"] = settings.steps;
688 root["cfg_scale"] = settings.cfg_scale;
689 root["batch_size"] = 1;
690 root["sampler_name"] = settings.sampler;
691 root["scheduler"] = settings.scheduler;
692 root["denoising_strength"] = settings.strength;
693 root["init_images"] = Json::arrayValue;
694 root["init_images"].append(encoded_input);
695 const std::string body = writeJson(root);
696 long status = 0;
697 const ResponseBuffer response = request(endpoint + "/sdapi/v1/img2img", &body, 3600L, status, &settings.cancelled);
698 const Json::Value document = parseJson(response.value, "sd-server");
699 if (status != 200) {
700 throw std::runtime_error("sd-server rejected frame: " + serverError(document));
701 }
702 if (!document["images"].isArray() || document["images"].empty() || !document["images"][0].isString()) {
703 throw std::runtime_error("sd-server response did not contain an output image");
704 }
705 encoded_output = document["images"][0].asString();
706 } else {
707 root["prompt"] = settings.prompt;
708 root["negative_prompt"] = settings.negative_prompt;
709 root["width"] = settings.width;
710 root["height"] = settings.height;
711 root["seed"] = settings.seed;
712 appendLoras(root, settings);
713 root["strength"] = settings.strength;
714 root["batch_count"] = 1;
715 root["init_image"] = encoded_input;
716 root["output_format"] = "png";
717 root["output_compression"] = 100;
718 Json::Value &sample = root["sample_params"];
719 sample["scheduler"] = settings.scheduler;
720 sample["sample_method"] = settings.sampler;
721 sample["sample_steps"] = settings.steps;
722 sample["guidance"]["txt_cfg"] = settings.cfg_scale;
723 Json::Value &hires = root["hires"];
724 const cv::Size working_size = neuralUpscaleWorkingSize(settings, rgba.size());
725 hires["enabled"] = true;
726 hires["upscaler"] = settings.upscale_model.stem().string();
727 hires["scale"] = 4.0;
728 hires["target_width"] = working_size.width;
729 hires["target_height"] = working_size.height;
730 hires["steps"] = settings.steps;
731 hires["denoising_strength"] = settings.strength;
732 hires["upscale_tile_size"] = 128;
733 encoded_output = submitAsyncJob(endpoint, "/sdcpp/v1/img_gen", root, settings);
734 }
735 const std::vector<std::uint8_t> decoded = decodeBase64(encoded_output);
736 cv::Mat encoded(1, static_cast<int>(decoded.size()), CV_8UC1, const_cast<std::uint8_t *>(decoded.data()));
737 cv::Mat output_bgr = cv::imdecode(encoded, cv::IMREAD_COLOR);
738 if (output_bgr.empty()) {
739 throw std::runtime_error("unable to decode image returned by sd-server");
740 }
741 cv::Mat output_rgba;
742 cv::cvtColor(output_bgr, output_rgba, cv::COLOR_BGR2RGBA);
743 if (!settings.upscale_model.empty() && settings.upscale_width > 0 && settings.upscale_height > 0 && output_rgba.size() != cv::Size(settings.upscale_width, settings.upscale_height)) {
744 cv::resize(output_rgba, output_rgba, {settings.upscale_width, settings.upscale_height}, 0.0, 0.0, cv::INTER_LANCZOS4);
745 } else if (settings.resize_to_input && output_rgba.size() != rgba.size()) {
746 cv::resize(output_rgba, output_rgba, rgba.size(), 0.0, 0.0, cv::INTER_LINEAR);
747 }
748 return output_rgba;
749 }
750} // namespace acmxvk::stable_diffusion
GLsizei GLsizei * length
cv::Mat process(const cv::Mat &rgba) const
std::filesystem::path diagnostic_log_path
void resetDiagnosticLog() const noexcept
std::size_t appendResponse(char *data, std::size_t size, std::size_t count, void *context)
cv::Size neuralUpscaleWorkingSize(const Settings &settings, const cv::Size &fallback)
std::string submitAsyncJob(const std::string &endpoint, std::string_view route, const Json::Value &request_root, const Settings &settings)
std::vector< std::uint8_t > decodeBase64(std::string_view encoded)
void appendLoras(Json::Value &root, const Settings &settings)
ResponseBuffer request(std::string_view url, const std::string *body, long timeout_seconds, long &status, const std::function< bool()> *cancelled)
int checkCancelled(void *context, curl_off_t, curl_off_t, curl_off_t, curl_off_t)
std::string encodeBase64(const std::vector< std::uint8_t > &bytes)
Json::Value parseJson(std::string_view text, std::string_view context)
char ** environ
std::vector< std::string > server_arguments