feat: add LoRA support

This commit is contained in:
leejet
2023-11-19 17:43:49 +08:00
parent 536f3af672
commit 9a9f3daf8e
7 changed files with 573 additions and 36 deletions

View File

@@ -268,6 +268,45 @@ struct ggml_tensor* ggml_group_norm_32(struct ggml_context* ctx,
return ggml_group_norm(ctx, a, 32);
}
std::pair<std::unordered_map<std::string, float>, std::string> extract_and_remove_lora(std::string text) {
std::regex re("<lora:([^:]+):([^>]+)>");
std::smatch matches;
std::unordered_map<std::string, float> filename2multiplier;
while (std::regex_search(text, matches, re)) {
std::string filename = matches[1].str();
float multiplier = std::stof(matches[2].str());
if (multiplier < 0.f) {
continue;
}
if (filename2multiplier.find(filename) == filename2multiplier.end()) {
filename2multiplier[filename] = multiplier;
} else {
filename2multiplier[filename] += multiplier;
}
text = std::regex_replace(text, re, "");
}
return std::make_pair(filename2multiplier, text);
}
bool ends_with(const std::string& str, const std::string& ending) {
if (str.length() >= ending.length()) {
return (str.compare(str.length() - ending.length(), ending.length(), ending) == 0);
} else {
return false;
}
}
void replace_all_chars(std::string& str, char target, char replacement) {
for (size_t i = 0; i < str.length(); ++i) {
if (str[i] == target) {
str[i] = replacement;
}
}
}
/*================================================== CLIPTokenizer ===================================================*/
const std::string UNK_TOKEN = "<|endoftext|>";
@@ -2794,6 +2833,16 @@ class StableDiffusionGGML {
UNetModel diffusion_model;
AutoEncoderKL first_stage_model;
std::map<std::string, struct ggml_tensor*> tensors;
std::string lora_model_dir;
// lora_name => lora_tensor_name => tensor
std::map<std::string, std::map<std::string, struct ggml_tensor*>> lora_tensors;
// lora_name => lora_params_ctx
std::map<std::string, ggml_context*> lora_params_ctxs;
// lora_name => multiplier
std::unordered_map<std::string, float> curr_lora_state;
std::shared_ptr<Denoiser> denoiser = std::make_shared<CompVisDenoiser>();
StableDiffusionGGML() = default;
@@ -2801,16 +2850,23 @@ class StableDiffusionGGML {
StableDiffusionGGML(int n_threads,
bool vae_decode_only,
bool free_params_immediately,
std::string lora_model_dir,
RNGType rng_type)
: n_threads(n_threads),
vae_decode_only(vae_decode_only),
free_params_immediately(free_params_immediately) {
free_params_immediately(free_params_immediately),
lora_model_dir(lora_model_dir) {
first_stage_model.decode_only = vae_decode_only;
if (rng_type == STD_DEFAULT_RNG) {
rng = std::make_shared<STDDefaultRNG>();
} else if (rng_type == CUDA_RNG) {
rng = std::make_shared<PhiloxRNG>();
}
if (lora_model_dir.size() > 0) {
if (lora_model_dir[lora_model_dir.size() - 1] != '/' && lora_model_dir[lora_model_dir.size() - 1] != '\\') {
this->lora_model_dir = lora_model_dir + "/";
}
}
}
~StableDiffusionGGML() {
@@ -2826,6 +2882,13 @@ class StableDiffusionGGML {
ggml_free(vae_params_ctx);
vae_params_ctx = NULL;
}
for (auto& kv : lora_params_ctxs) {
ggml_free(kv.second);
}
lora_params_ctxs.clear();
tensors.clear();
lora_tensors.clear();
}
bool load_from_file(const std::string& file_path, Schedule schedule) {
@@ -2963,8 +3026,6 @@ class StableDiffusionGGML {
}
}
std::map<std::string, struct ggml_tensor*> tensors;
LOG_DEBUG("preparing memory for the weights");
// prepare memory for the weights
{
@@ -3255,6 +3316,306 @@ class StableDiffusionGGML {
return result < -1;
}
bool load_lora_from_file(const std::string& lora_name) {
if (lora_tensors.find(lora_name) != lora_tensors.end()) {
return true;
}
std::string file_path = lora_model_dir + lora_name + "-ggml-lora.bin";
LOG_INFO("loading lora '%s' from '%s'", lora_name.c_str(), file_path.c_str());
std::ifstream file(file_path, std::ios::binary);
if (!file.is_open()) {
LOG_ERROR("failed to open '%s'", file_path.c_str());
return false;
}
// get file size
file.seekg(0, file.end);
int file_size = (int)file.tellg();
file.seekg(0, file.beg);
LOG_DEBUG("'%s': %.2fMB", file_path.c_str(), file_size * 1.f / 1024 / 1024);
LOG_DEBUG("verifying magic");
// verify magic
{
uint32_t magic;
file.read(reinterpret_cast<char*>(&magic), sizeof(magic));
if (magic != GGML_FILE_MAGIC) {
LOG_ERROR("invalid model file '%s' (bad magic)", file_path.c_str());
return false;
}
}
LOG_DEBUG("loading hparams");
// load hparams
file.read(reinterpret_cast<char*>(&ftype), sizeof(ftype));
int model_type = (ftype >> 16) & 0xFFFF;
if (model_type >= MODEL_TYPE_COUNT) {
LOG_ERROR("invalid model file '%s' (bad model type value %d)", file_path.c_str(), ftype);
return false;
}
LOG_INFO("lora model type: %s", model_type_to_str[model_type]);
ggml_type wtype = ggml_ftype_to_ggml_type((ggml_ftype)(ftype & 0xFFFF));
LOG_INFO("ftype: %s", ggml_type_name(wtype));
if (wtype == GGML_TYPE_COUNT) {
LOG_ERROR("invalid model file '%s' (bad ftype value %d)", file_path.c_str(), ftype);
return false;
}
// create the ggml context for network params
struct ggml_init_params params;
size_t ctx_size = 10 * 1024 * 1024; // 10 MB, for padding
ctx_size += file_size;
params.mem_size = ctx_size;
params.mem_buffer = NULL;
params.no_alloc = false;
params.dynamic = false;
LOG_DEBUG("lora '%s' params ctx size = % 6.2f MB", lora_name.c_str(), ctx_size / (1024.0 * 1024.0));
ggml_context* lora_params_ctx = ggml_init(params);
if (!lora_params_ctx) {
LOG_ERROR("ggml_init() failed");
return false;
}
lora_params_ctxs[lora_name] = lora_params_ctx;
std::map<std::string, struct ggml_tensor*> lora_tensor_map;
int64_t t0 = ggml_time_ms();
// load weights
{
int n_tensors = 0;
size_t total_size = 0;
while (true) {
int32_t n_dims;
int32_t length;
int32_t ttype;
file.read(reinterpret_cast<char*>(&n_dims), sizeof(n_dims));
file.read(reinterpret_cast<char*>(&length), sizeof(length));
file.read(reinterpret_cast<char*>(&ttype), sizeof(ttype));
if (file.eof()) {
break;
}
int32_t nelements = 1;
int32_t ne[4] = {1, 1, 1, 1};
for (int i = 0; i < n_dims; ++i) {
file.read(reinterpret_cast<char*>(&ne[i]), sizeof(ne[i]));
nelements *= ne[i];
}
const size_t num_bytes = nelements / ggml_blck_size(ggml_type(ttype)) * ggml_type_size(ggml_type(ttype));
std::string name_buf(length, 0);
file.read(&name_buf[0], length);
std::string name = std::string(name_buf.data());
// LOG_DEBUG("load lora tensor %s", name.c_str());
int64_t ne64[4] = {ne[0], ne[1], ne[2], ne[3]};
struct ggml_tensor* tensor = ggml_new_tensor(lora_params_ctx, (ggml_type)ttype, n_dims, ne64);
file.read(reinterpret_cast<char*>(tensor->data), num_bytes);
lora_tensor_map[name] = tensor;
total_size += ggml_nbytes(tensor);
}
}
lora_tensors[lora_name] = lora_tensor_map;
int64_t t1 = ggml_time_ms();
LOG_INFO("lora '%s' params size = %.2fMB",
lora_name.c_str(),
ggml_used_mem(lora_params_ctx) / 1024.0 / 1024.0);
LOG_INFO("loading lora from '%s' completed, taking %.2fs", file_path.c_str(), (t1 - t0) * 1.0f / 1000);
file.close();
return true;
}
void remove_lora_params(const std::string& lora_name) {
if (lora_params_ctxs.find(lora_name) == lora_params_ctxs.end()) {
return;
}
ggml_free(lora_params_ctxs[lora_name]);
lora_params_ctxs.erase(lora_name);
lora_tensors.erase(lora_name);
}
void apply_lora(const std::string& lora_name, float multiplier) {
int64_t t0 = ggml_time_ms();
if (!load_lora_from_file(lora_name)) {
std::string file_path = lora_model_dir + lora_name + "-ggml-lora.bin";
LOG_WARN("apply lora '%s' failed", lora_name.c_str());
return;
}
size_t ctx_size = 500 * 1024 * 1024; // 500MB
void* mem_buffer = malloc(ctx_size);
if (!mem_buffer) {
if (free_params_immediately) {
remove_lora_params(lora_name);
}
LOG_ERROR("malloc() failed");
return;
}
std::map<std::string, struct ggml_tensor*>& lora_tensor_map = lora_tensors[lora_name];
std::set<std::string> applied_lora_tensors;
for (auto& kv : tensors) {
const std::string name = kv.first;
ggml_tensor* weight = kv.second;
std::string ending = ".weight";
if (!ends_with(name, ending)) {
continue;
}
// find corresponding lora tensors
std::string network_name = name.substr(0, name.size() - ending.size()); // remove .weight
replace_all_chars(network_name, '.', '_');
std::string lora_up_name = network_name + ".lora_up.weight";
std::string lora_down_name = network_name + ".lora_down.weight";
std::string alpha_name = network_name + ".alpha";
std::string scale_name = network_name + ".scale";
ggml_tensor* lora_up = NULL;
ggml_tensor* lora_down = NULL;
float scale = 1.0f;
if (lora_tensor_map.find(lora_up_name) != lora_tensor_map.end()) {
lora_up = lora_tensor_map[lora_up_name];
}
if (lora_tensor_map.find(lora_down_name) != lora_tensor_map.end()) {
lora_down = lora_tensor_map[lora_down_name];
}
if (lora_up == NULL || lora_down == NULL) {
continue;
}
// LOG_DEBUG("apply lora tensor %s [%ld %ld %ld %ld]", network_name.c_str(), weight->ne[0], weight->ne[1], weight->ne[2], weight->ne[3]);
applied_lora_tensors.insert(lora_up_name);
applied_lora_tensors.insert(lora_down_name);
applied_lora_tensors.insert(alpha_name);
applied_lora_tensors.insert(scale_name);
// calc_scale
int64_t dim = lora_down->ne[lora_down->n_dims - 1];
if (lora_tensor_map.find(scale_name) != lora_tensor_map.end()) {
ggml_tensor* t = lora_tensor_map[scale_name];
scale = ggml_get_f32_1d(t, 0);
} else if (lora_tensor_map.find(alpha_name) != lora_tensor_map.end()) {
ggml_tensor* t = lora_tensor_map[alpha_name];
scale = ggml_get_f32_1d(t, 0) / dim;
}
// LOG_DEBUG("scale: %f %ld", scale, dim);
scale = scale * multiplier;
// apply
{
struct ggml_init_params params;
params.mem_size = ctx_size;
params.mem_buffer = mem_buffer;
params.no_alloc = false;
params.dynamic = false;
struct ggml_context* ctx = ggml_init(params);
if (!ctx) {
LOG_ERROR("ggml_init() failed");
free(mem_buffer);
if (free_params_immediately) {
remove_lora_params(lora_name);
}
return;
}
ggml_tensor* scale_factor = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1);
ggml_set_f32_1d(scale_factor, 0, scale);
int64_t lora_up_size_0 = lora_up->ne[lora_up->n_dims - 1];
lora_up = ggml_reshape_2d(ctx, lora_up, ggml_nelements(lora_up) / lora_up_size_0, lora_up_size_0);
int64_t lora_down_size_0 = lora_down->ne[lora_down->n_dims - 1];
lora_down = ggml_reshape_2d(ctx, lora_down, ggml_nelements(lora_down) / lora_down_size_0, lora_down_size_0);
lora_down = ggml_cont(ctx, ggml_transpose(ctx, lora_down));
if (lora_down->type != GGML_TYPE_F32) {
ggml_tensor* lora_down_f32 = ggml_new_tensor(ctx, GGML_TYPE_F32, lora_down->n_dims, lora_down->ne);
lora_down = ggml_cpy_inplace(ctx, lora_down, lora_down_f32);
}
ggml_tensor* updown = ggml_mul_mat(ctx, lora_up, lora_down);
updown = ggml_cont(ctx, ggml_transpose(ctx, updown));
updown = ggml_reshape(ctx, updown, weight);
GGML_ASSERT(ggml_nelements(updown) == ggml_nelements(weight));
updown = ggml_scale_inplace(ctx, updown, scale_factor);
ggml_tensor* final_weight;
final_weight = ggml_add_inplace(ctx, weight, updown);
final_weight = ggml_cpy_inplace(ctx, final_weight, weight);
struct ggml_cgraph* graph = ggml_build_forward_ctx(ctx, final_weight);
ggml_graph_compute_with_ctx(ctx, graph, n_threads);
// LOG_INFO("network_name '%s' ggml_used_mem size = %.2fMB",
// network_name.c_str(),
// ggml_used_mem(ctx) / 1024.0 / 1024.0);
ggml_free(ctx);
}
}
free(mem_buffer);
for (auto& kv : lora_tensor_map) {
if (applied_lora_tensors.find(kv.first) == applied_lora_tensors.end()) {
LOG_WARN("unused lora tensor %s", kv.first.c_str());
}
}
if (free_params_immediately) {
remove_lora_params(lora_name);
}
int64_t t1 = ggml_time_ms();
LOG_INFO("apply lora '%s:%f' completed, taking %.2fs",
lora_name.c_str(),
multiplier,
(t1 - t0) * 1.0f / 1000);
}
void apply_loras(const std::unordered_map<std::string, float>& lora_state) {
std::unordered_map<std::string, float> lora_state_diff;
for (auto& kv : lora_state) {
const std::string& lora_name = kv.first;
float multiplier = kv.second;
if (curr_lora_state.find(lora_name) != curr_lora_state.end()) {
float curr_multiplier = curr_lora_state[lora_name];
float multiplier_diff = multiplier - curr_multiplier;
if (multiplier_diff != 0.f) {
lora_state_diff[lora_name] = multiplier_diff;
}
} else {
lora_state_diff[lora_name] = multiplier;
}
}
for (auto& kv : lora_state_diff) {
apply_lora(kv.first, kv.second);
}
curr_lora_state = lora_state;
}
ggml_tensor* get_learned_condition(ggml_context* res_ctx, const std::string& text) {
auto tokens_and_weights = cond_stage_model.tokenize(text,
cond_stage_model.text_model.max_position_embeddings,
@@ -4235,10 +4596,12 @@ class StableDiffusionGGML {
StableDiffusion::StableDiffusion(int n_threads,
bool vae_decode_only,
bool free_params_immediately,
std::string lora_model_dir,
RNGType rng_type) {
sd = std::make_shared<StableDiffusionGGML>(n_threads,
vae_decode_only,
free_params_immediately,
lora_model_dir,
rng_type);
}
@@ -4246,8 +4609,8 @@ bool StableDiffusion::load_from_file(const std::string& file_path, Schedule s) {
return sd->load_from_file(file_path, s);
}
std::vector<uint8_t> StableDiffusion::txt2img(const std::string& prompt,
const std::string& negative_prompt,
std::vector<uint8_t> StableDiffusion::txt2img(std::string prompt,
std::string negative_prompt,
float cfg_scale,
int width,
int height,
@@ -4272,13 +4635,28 @@ std::vector<uint8_t> StableDiffusion::txt2img(const std::string& prompt,
}
sd->rng->manual_seed(seed);
// extract and remote lora
auto result_pair = extract_and_remove_lora(prompt);
std::unordered_map<std::string, float> lora_f2m = result_pair.first; // lora_name -> multiplier
for (auto& kv : lora_f2m) {
LOG_DEBUG("lora %s:%.2f", kv.first.c_str(), kv.second);
}
prompt = result_pair.second;
LOG_DEBUG("prompt after extract and remote lora: \"%s\"", prompt.c_str());
// load lora from file
int64_t t0 = ggml_time_ms();
sd->apply_loras(lora_f2m);
int64_t t1 = ggml_time_ms();
LOG_INFO("apply_loras completed, taking %.2fs", (t1 - t0) * 1.0f / 1000);
t0 = ggml_time_ms();
ggml_tensor* c = sd->get_learned_condition(ctx, prompt);
struct ggml_tensor* uc = NULL;
if (cfg_scale != 1.0) {
uc = sd->get_learned_condition(ctx, negative_prompt);
}
int64_t t1 = ggml_time_ms();
t1 = ggml_time_ms();
LOG_INFO("get_learned_condition completed, taking %.2fs", (t1 - t0) * 1.0f / 1000);
if (sd->free_params_immediately) {
@@ -4334,8 +4712,8 @@ std::vector<uint8_t> StableDiffusion::txt2img(const std::string& prompt,
}
std::vector<uint8_t> StableDiffusion::img2img(const std::vector<uint8_t>& init_img_vec,
const std::string& prompt,
const std::string& negative_prompt,
std::string prompt,
std::string negative_prompt,
float cfg_scale,
int width,
int height,
@@ -4372,14 +4750,29 @@ std::vector<uint8_t> StableDiffusion::img2img(const std::vector<uint8_t>& init_i
}
sd->rng->manual_seed(seed);
// extract and remote lora
auto result_pair = extract_and_remove_lora(prompt);
std::unordered_map<std::string, float> lora_f2m = result_pair.first; // lora_name -> multiplier
for (auto& kv : lora_f2m) {
LOG_DEBUG("lora %s:%.2f", kv.first.c_str(), kv.second);
}
prompt = result_pair.second;
LOG_DEBUG("prompt after extract and remote lora: \"%s\"", prompt.c_str());
// load lora from file
int64_t t0 = ggml_time_ms();
sd->apply_loras(lora_f2m);
int64_t t1 = ggml_time_ms();
LOG_INFO("apply_loras completed, taking %.2fs", (t1 - t0) * 1.0f / 1000);
ggml_tensor* init_img = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, width, height, 3, 1);
image_vec_to_ggml(init_img_vec, init_img);
int64_t t0 = ggml_time_ms();
t0 = ggml_time_ms();
ggml_tensor* moments = sd->encode_first_stage(ctx, init_img);
ggml_tensor* init_latent = sd->get_first_stage_encoding(ctx, moments);
// print_ggml_tensor(init_latent);
int64_t t1 = ggml_time_ms();
t1 = ggml_time_ms();
LOG_INFO("encode_first_stage completed, taking %.2fs", (t1 - t0) * 1.0f / 1000);
ggml_reset_curr_max_dynamic_size(); // reset counter