mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-08-10 06:07:58 +00:00
fix: skip incompatible LoRA weights (#1825)
This commit is contained in:
parent
5ef4a7557d
commit
22516991cb
@ -14,6 +14,8 @@ struct LoraModel : public GGMLRunner {
|
|||||||
std::unordered_map<std::string, ggml_tensor*> lora_tensors;
|
std::unordered_map<std::string, ggml_tensor*> lora_tensors;
|
||||||
std::map<ggml_tensor*, ggml_tensor*> original_tensor_to_final_tensor;
|
std::map<ggml_tensor*, ggml_tensor*> original_tensor_to_final_tensor;
|
||||||
std::set<std::string> applied_lora_tensors;
|
std::set<std::string> applied_lora_tensors;
|
||||||
|
std::set<std::string> skipped_incompatible_lora_tensors;
|
||||||
|
std::set<std::string> warned_incompatible_model_tensors;
|
||||||
std::string file_path;
|
std::string file_path;
|
||||||
std::shared_ptr<ModelManager> model_manager;
|
std::shared_ptr<ModelManager> model_manager;
|
||||||
ggml_backend_t params_backend = nullptr;
|
ggml_backend_t params_backend = nullptr;
|
||||||
@ -133,6 +135,8 @@ struct LoraModel : public GGMLRunner {
|
|||||||
lora_tensors.clear();
|
lora_tensors.clear();
|
||||||
original_tensor_to_final_tensor.clear();
|
original_tensor_to_final_tensor.clear();
|
||||||
applied_lora_tensors.clear();
|
applied_lora_tensors.clear();
|
||||||
|
skipped_incompatible_lora_tensors.clear();
|
||||||
|
warned_incompatible_model_tensors.clear();
|
||||||
applied = false;
|
applied = false;
|
||||||
tensor_preprocessed = false;
|
tensor_preprocessed = false;
|
||||||
}
|
}
|
||||||
@ -546,7 +550,27 @@ struct LoraModel : public GGMLRunner {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
GGML_ASSERT(ggml_nelements(diff) == ggml_nelements(model_tensor));
|
if (ggml_nelements(diff) != ggml_nelements(model_tensor)) {
|
||||||
|
const std::string lora_tensor_prefix = "lora." + model_tensor_name + ".";
|
||||||
|
for (const auto& tensor_name : applied_lora_tensors) {
|
||||||
|
if (starts_with(tensor_name, lora_tensor_prefix)) {
|
||||||
|
skipped_incompatible_lora_tensors.insert(tensor_name);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if (warned_incompatible_model_tensors.insert(model_tensor_name).second) {
|
||||||
|
LOG_WARN("skip incompatible LoRA tensor |%s|: model shape = [%lld, %lld, %lld, %lld], LoRA shape = [%lld, %lld, %lld, %lld]",
|
||||||
|
model_tensor_name.c_str(),
|
||||||
|
static_cast<long long>(model_tensor->ne[0]),
|
||||||
|
static_cast<long long>(model_tensor->ne[1]),
|
||||||
|
static_cast<long long>(model_tensor->ne[2]),
|
||||||
|
static_cast<long long>(model_tensor->ne[3]),
|
||||||
|
static_cast<long long>(diff->ne[0]),
|
||||||
|
static_cast<long long>(diff->ne[1]),
|
||||||
|
static_cast<long long>(diff->ne[2]),
|
||||||
|
static_cast<long long>(diff->ne[3]));
|
||||||
|
}
|
||||||
|
return nullptr;
|
||||||
|
}
|
||||||
diff = ggml_reshape(ctx, diff, model_tensor);
|
diff = ggml_reshape(ctx, diff, model_tensor);
|
||||||
}
|
}
|
||||||
return diff;
|
return diff;
|
||||||
@ -555,6 +579,7 @@ struct LoraModel : public GGMLRunner {
|
|||||||
ggml_tensor* get_out_diff(ggml_context* ctx,
|
ggml_tensor* get_out_diff(ggml_context* ctx,
|
||||||
ggml_backend_t backend,
|
ggml_backend_t backend,
|
||||||
ggml_tensor* x,
|
ggml_tensor* x,
|
||||||
|
ggml_tensor* model_weight,
|
||||||
WeightAdapter::ForwardParams forward_params,
|
WeightAdapter::ForwardParams forward_params,
|
||||||
const std::string& model_tensor_name) {
|
const std::string& model_tensor_name) {
|
||||||
ggml_tensor* out_diff = nullptr;
|
ggml_tensor* out_diff = nullptr;
|
||||||
@ -707,6 +732,43 @@ struct LoraModel : public GGMLRunner {
|
|||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if (!is_conv2d) {
|
||||||
|
const int64_t down_in = lora_down->ne[0];
|
||||||
|
const int64_t down_out = lora_down->ne[1];
|
||||||
|
const int64_t up_in = lora_up->ne[0];
|
||||||
|
const int64_t up_out = lora_up->ne[1];
|
||||||
|
|
||||||
|
bool compatible = down_in == model_weight->ne[0] &&
|
||||||
|
up_out == model_weight->ne[1];
|
||||||
|
if (lora_mid != nullptr) {
|
||||||
|
compatible = compatible &&
|
||||||
|
lora_mid->ne[0] == down_out &&
|
||||||
|
up_in == lora_mid->ne[1];
|
||||||
|
} else {
|
||||||
|
compatible = compatible && up_in == down_out;
|
||||||
|
}
|
||||||
|
|
||||||
|
if (!compatible) {
|
||||||
|
skipped_incompatible_lora_tensors.insert(lora_down_name);
|
||||||
|
skipped_incompatible_lora_tensors.insert(lora_up_name);
|
||||||
|
skipped_incompatible_lora_tensors.insert(lora_mid_name);
|
||||||
|
skipped_incompatible_lora_tensors.insert(scale_name);
|
||||||
|
skipped_incompatible_lora_tensors.insert(alpha_name);
|
||||||
|
if (warned_incompatible_model_tensors.insert(model_tensor_name).second) {
|
||||||
|
LOG_WARN("skip incompatible LoRA tensor |%s|: model shape = [%lld, %lld], down shape = [%lld, %lld], up shape = [%lld, %lld]",
|
||||||
|
model_tensor_name.c_str(),
|
||||||
|
static_cast<long long>(model_weight->ne[0]),
|
||||||
|
static_cast<long long>(model_weight->ne[1]),
|
||||||
|
static_cast<long long>(down_in),
|
||||||
|
static_cast<long long>(down_out),
|
||||||
|
static_cast<long long>(up_in),
|
||||||
|
static_cast<long long>(up_out));
|
||||||
|
}
|
||||||
|
index++;
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
applied_lora_tensors.insert(lora_up_name);
|
applied_lora_tensors.insert(lora_up_name);
|
||||||
applied_lora_tensors.insert(lora_down_name);
|
applied_lora_tensors.insert(lora_down_name);
|
||||||
|
|
||||||
@ -869,10 +931,13 @@ struct LoraModel : public GGMLRunner {
|
|||||||
void stat(bool at_runntime = false) {
|
void stat(bool at_runntime = false) {
|
||||||
size_t total_lora_tensors_count = 0;
|
size_t total_lora_tensors_count = 0;
|
||||||
size_t applied_lora_tensors_count = 0;
|
size_t applied_lora_tensors_count = 0;
|
||||||
|
size_t skipped_lora_tensors_count = 0;
|
||||||
|
|
||||||
for (auto& kv : lora_tensors) {
|
for (auto& kv : lora_tensors) {
|
||||||
total_lora_tensors_count++;
|
total_lora_tensors_count++;
|
||||||
if (applied_lora_tensors.find(kv.first) == applied_lora_tensors.end()) {
|
if (skipped_incompatible_lora_tensors.find(kv.first) != skipped_incompatible_lora_tensors.end()) {
|
||||||
|
skipped_lora_tensors_count++;
|
||||||
|
} else if (applied_lora_tensors.find(kv.first) == applied_lora_tensors.end()) {
|
||||||
if (!at_runntime) {
|
if (!at_runntime) {
|
||||||
LOG_WARN("unused lora tensor |%s|", kv.first.c_str());
|
LOG_WARN("unused lora tensor |%s|", kv.first.c_str());
|
||||||
print_ggml_tensor(kv.second, true);
|
print_ggml_tensor(kv.second, true);
|
||||||
@ -884,12 +949,17 @@ struct LoraModel : public GGMLRunner {
|
|||||||
/* Don't worry if this message shows up twice in the logs per LoRA,
|
/* Don't worry if this message shows up twice in the logs per LoRA,
|
||||||
* this function is called once to calculate the required buffer size
|
* this function is called once to calculate the required buffer size
|
||||||
* and then again to actually generate a graph to be used */
|
* and then again to actually generate a graph to be used */
|
||||||
if (!at_runntime && applied_lora_tensors_count != total_lora_tensors_count) {
|
size_t compatible_lora_tensors_count = total_lora_tensors_count - skipped_lora_tensors_count;
|
||||||
|
if (!at_runntime && applied_lora_tensors_count != compatible_lora_tensors_count) {
|
||||||
LOG_WARN("Only (%lu / %lu) LoRA tensors have been applied, lora_file_path = %s",
|
LOG_WARN("Only (%lu / %lu) LoRA tensors have been applied, lora_file_path = %s",
|
||||||
applied_lora_tensors_count, total_lora_tensors_count, file_path.c_str());
|
applied_lora_tensors_count, compatible_lora_tensors_count, file_path.c_str());
|
||||||
} else {
|
} else {
|
||||||
LOG_INFO("(%lu / %lu) LoRA tensors have been applied, lora_file_path = %s",
|
LOG_INFO("(%lu / %lu) LoRA tensors have been applied, lora_file_path = %s",
|
||||||
applied_lora_tensors_count, total_lora_tensors_count, file_path.c_str());
|
applied_lora_tensors_count, compatible_lora_tensors_count, file_path.c_str());
|
||||||
|
}
|
||||||
|
if (skipped_lora_tensors_count > 0) {
|
||||||
|
LOG_WARN("(%lu / %lu) incompatible LoRA tensors have been skipped, lora_file_path = %s",
|
||||||
|
skipped_lora_tensors_count, total_lora_tensors_count, file_path.c_str());
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
};
|
};
|
||||||
@ -953,7 +1023,7 @@ public:
|
|||||||
forward_params.conv2d.scale);
|
forward_params.conv2d.scale);
|
||||||
}
|
}
|
||||||
for (auto& lora_model : lora_models) {
|
for (auto& lora_model : lora_models) {
|
||||||
ggml_tensor* out_diff = lora_model->get_out_diff(ctx, backend, x, forward_params, prefix + "weight");
|
ggml_tensor* out_diff = lora_model->get_out_diff(ctx, backend, x, w, forward_params, prefix + "weight");
|
||||||
if (out_diff == nullptr) {
|
if (out_diff == nullptr) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user