diff --git a/src/model/adapter/lora.hpp b/src/model/adapter/lora.hpp index 0b759175b..b5a84c2ef 100644 --- a/src/model/adapter/lora.hpp +++ b/src/model/adapter/lora.hpp @@ -14,6 +14,8 @@ struct LoraModel : public GGMLRunner { std::unordered_map lora_tensors; std::map original_tensor_to_final_tensor; std::set applied_lora_tensors; + std::set skipped_incompatible_lora_tensors; + std::set warned_incompatible_model_tensors; std::string file_path; std::shared_ptr model_manager; ggml_backend_t params_backend = nullptr; @@ -133,6 +135,8 @@ struct LoraModel : public GGMLRunner { lora_tensors.clear(); original_tensor_to_final_tensor.clear(); applied_lora_tensors.clear(); + skipped_incompatible_lora_tensors.clear(); + warned_incompatible_model_tensors.clear(); applied = 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(model_tensor->ne[0]), + static_cast(model_tensor->ne[1]), + static_cast(model_tensor->ne[2]), + static_cast(model_tensor->ne[3]), + static_cast(diff->ne[0]), + static_cast(diff->ne[1]), + static_cast(diff->ne[2]), + static_cast(diff->ne[3])); + } + return nullptr; + } diff = ggml_reshape(ctx, diff, model_tensor); } return diff; @@ -555,6 +579,7 @@ struct LoraModel : public GGMLRunner { ggml_tensor* get_out_diff(ggml_context* ctx, ggml_backend_t backend, ggml_tensor* x, + ggml_tensor* model_weight, WeightAdapter::ForwardParams forward_params, const std::string& model_tensor_name) { ggml_tensor* out_diff = nullptr; @@ -707,6 +732,43 @@ struct LoraModel : public GGMLRunner { 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(model_weight->ne[0]), + static_cast(model_weight->ne[1]), + static_cast(down_in), + static_cast(down_out), + static_cast(up_in), + static_cast(up_out)); + } + index++; + continue; + } + } + applied_lora_tensors.insert(lora_up_name); applied_lora_tensors.insert(lora_down_name); @@ -869,10 +931,13 @@ struct LoraModel : public GGMLRunner { void stat(bool at_runntime = false) { size_t total_lora_tensors_count = 0; size_t applied_lora_tensors_count = 0; + size_t skipped_lora_tensors_count = 0; for (auto& kv : lora_tensors) { 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) { LOG_WARN("unused lora tensor |%s|", kv.first.c_str()); 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, * this function is called once to calculate the required buffer size * 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", - 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 { 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 @@ struct MultiLoraAdapter : public WeightAdapter { forward_params.conv2d.scale); } 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) { continue; }