mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-08-11 06:36:39 +00:00
add euler flow flash sample method
This commit is contained in:
parent
2f89058f24
commit
e56295d180
@ -37,6 +37,7 @@ enum rng_type_t {
|
|||||||
|
|
||||||
enum sample_method_t {
|
enum sample_method_t {
|
||||||
EULER_SAMPLE_METHOD,
|
EULER_SAMPLE_METHOD,
|
||||||
|
EULER_FLOW_FLASH_SAMPLE_METHOD,
|
||||||
EULER_A_SAMPLE_METHOD,
|
EULER_A_SAMPLE_METHOD,
|
||||||
HEUN_SAMPLE_METHOD,
|
HEUN_SAMPLE_METHOD,
|
||||||
DPM2_SAMPLE_METHOD,
|
DPM2_SAMPLE_METHOD,
|
||||||
|
|||||||
@ -867,6 +867,31 @@ static sd::Tensor<float> sample_euler_flow(denoise_cb_t model,
|
|||||||
return x;
|
return x;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
static sd::Tensor<float> sample_euler_flow_flash(denoise_cb_t model,
|
||||||
|
sd::Tensor<float> x,
|
||||||
|
const std::vector<float>& sigmas,
|
||||||
|
std::shared_ptr<RNG> rng,
|
||||||
|
float eta) {
|
||||||
|
float s_noise = eta;
|
||||||
|
int steps = static_cast<int>(sigmas.size()) - 1;
|
||||||
|
for (int i = 0; i < steps; i++) {
|
||||||
|
float sigma = sigmas[i];
|
||||||
|
float sigma_next = sigmas[i + 1];
|
||||||
|
auto denoised_opt = model(x, sigma, i + 1);
|
||||||
|
if (denoised_opt.empty()) {
|
||||||
|
return {};
|
||||||
|
}
|
||||||
|
sd::Tensor<float> denoised = std::move(denoised_opt);
|
||||||
|
if (sigma_next == 0.0f) {
|
||||||
|
x = std::move(denoised);
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
auto noise = sd::Tensor<float>::randn_like(x, rng);
|
||||||
|
x = sigma_next * noise * s_noise + (1.0f - sigma_next) * denoised;
|
||||||
|
}
|
||||||
|
return x;
|
||||||
|
}
|
||||||
|
|
||||||
static sd::Tensor<float> sample_euler(denoise_cb_t model,
|
static sd::Tensor<float> sample_euler(denoise_cb_t model,
|
||||||
sd::Tensor<float> x,
|
sd::Tensor<float> x,
|
||||||
const std::vector<float>& sigmas) {
|
const std::vector<float>& sigmas) {
|
||||||
@ -1658,6 +1683,8 @@ static sd::Tensor<float> sample_k_diffusion(sample_method_t method,
|
|||||||
float eta,
|
float eta,
|
||||||
bool is_flow_denoiser) {
|
bool is_flow_denoiser) {
|
||||||
switch (method) {
|
switch (method) {
|
||||||
|
case EULER_FLOW_FLASH_SAMPLE_METHOD:
|
||||||
|
return sample_euler_flow_flash(model, std::move(x), sigmas, rng, eta);
|
||||||
case EULER_A_SAMPLE_METHOD:
|
case EULER_A_SAMPLE_METHOD:
|
||||||
if (is_flow_denoiser)
|
if (is_flow_denoiser)
|
||||||
return sample_euler_flow(model, std::move(x), sigmas, rng, eta);
|
return sample_euler_flow(model, std::move(x), sigmas, rng, eta);
|
||||||
|
|||||||
@ -60,6 +60,7 @@ const char* model_version_to_str[] = {
|
|||||||
|
|
||||||
const char* sampling_methods_str[] = {
|
const char* sampling_methods_str[] = {
|
||||||
"Euler",
|
"Euler",
|
||||||
|
"Euler Flow Flash",
|
||||||
"Euler A",
|
"Euler A",
|
||||||
"Heun",
|
"Heun",
|
||||||
"DPM2",
|
"DPM2",
|
||||||
@ -1978,6 +1979,7 @@ enum rng_type_t str_to_rng_type(const char* str) {
|
|||||||
|
|
||||||
const char* sample_method_to_str[] = {
|
const char* sample_method_to_str[] = {
|
||||||
"euler",
|
"euler",
|
||||||
|
"euler_flow_flash",
|
||||||
"euler_a",
|
"euler_a",
|
||||||
"heun",
|
"heun",
|
||||||
"dpm2",
|
"dpm2",
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user