fix: skip incompatible LoRA weights (#1825)

This commit is contained in:
leejet 2026-07-28 00:06:46 +08:00 committed by GitHub
parent 5ef4a7557d
commit 22516991cb
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -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;
} }