mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-08-01 07:40:41 -05:00
add a get_learned_condition_with_trigger function to do photomaker stuff
This commit is contained in:
43
clip.hpp
43
clip.hpp
@@ -76,8 +76,11 @@ private:
|
||||
SDVersion version = VERSION_1_x;
|
||||
std::map<int, std::u32string> byte_encoder;
|
||||
std::map<std::u32string, int> encoder;
|
||||
std::map<int, std::u32string> decoder;
|
||||
std::map<std::pair<std::u32string, std::u32string>, int> bpe_ranks;
|
||||
std::regex pat;
|
||||
int encoder_len;
|
||||
int bpe_len;
|
||||
|
||||
static std::string strip(const std::string& str) {
|
||||
std::string::size_type start = str.find_first_not_of(" \t\n\r\v\f");
|
||||
@@ -118,6 +121,7 @@ public:
|
||||
|
||||
void load_from_merges(const std::string& merges_utf8_str) {
|
||||
auto byte_unicode_pairs = bytes_to_unicode();
|
||||
// printf("byte_unicode_pairs have %lu pairs \n", byte_unicode_pairs.size());
|
||||
byte_encoder = std::map<int, std::u32string>(byte_unicode_pairs.begin(), byte_unicode_pairs.end());
|
||||
// for (auto & pair: byte_unicode_pairs) {
|
||||
// std::cout << pair.first << ": " << pair.second << std::endl;
|
||||
@@ -138,6 +142,8 @@ public:
|
||||
size_t space_pos = merge.find(' ');
|
||||
merge_pairs.emplace_back(merge.substr(0, space_pos), merge.substr(space_pos + 1));
|
||||
// LOG_DEBUG("%s", utf32_to_utf8(merge.substr(space_pos + 1)).c_str());
|
||||
// printf("%s :: %s, %s \n", utf32_to_utf8(merge).c_str(), utf32_to_utf8(merge.substr(0, space_pos)).c_str(),
|
||||
// utf32_to_utf8(merge.substr(space_pos + 1)).c_str());
|
||||
}
|
||||
std::vector<std::u32string> vocab;
|
||||
for (const auto& pair : byte_unicode_pairs) {
|
||||
@@ -154,15 +160,37 @@ public:
|
||||
LOG_DEBUG("vocab size: %llu", vocab.size());
|
||||
int i = 0;
|
||||
for (const auto& token : vocab) {
|
||||
encoder[token] = i++;
|
||||
encoder[token] = i;
|
||||
decoder[i] = token;
|
||||
i++;
|
||||
}
|
||||
encoder_len = i;
|
||||
|
||||
auto it = encoder.find(utf8_to_utf32("img"));
|
||||
if (it != encoder.end()) {
|
||||
printf(" trigger word img already in vocab \n");
|
||||
}else{
|
||||
printf(" trigger word img not in vocab yet\n");
|
||||
}
|
||||
|
||||
int rank = 0;
|
||||
for (const auto& merge : merge_pairs) {
|
||||
bpe_ranks[merge] = rank++;
|
||||
}
|
||||
bpe_len = rank;
|
||||
};
|
||||
|
||||
|
||||
void add_token(const std::string &text){
|
||||
std::u32string token = utf8_to_utf32(text);
|
||||
auto it = encoder.find(token);
|
||||
if (it != encoder.end()) {
|
||||
encoder[token] = encoder_len;
|
||||
decoder[encoder_len] = token;
|
||||
encoder_len++;
|
||||
}
|
||||
}
|
||||
|
||||
std::u32string bpe(const std::u32string& token) {
|
||||
std::vector<std::u32string> word;
|
||||
|
||||
@@ -243,6 +271,7 @@ public:
|
||||
size_t max_length = 0,
|
||||
bool padding = false) {
|
||||
std::vector<int32_t> tokens = encode(text, on_new_token_cb);
|
||||
|
||||
tokens.insert(tokens.begin(), BOS_TOKEN_ID);
|
||||
if (max_length > 0) {
|
||||
if (tokens.size() > max_length - 1) {
|
||||
@@ -259,6 +288,7 @@ public:
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return tokens;
|
||||
}
|
||||
|
||||
@@ -308,7 +338,8 @@ public:
|
||||
ss << "\"" << token << "\", ";
|
||||
}
|
||||
ss << "]";
|
||||
LOG_DEBUG("split prompt \"%s\" to tokens %s", original_text.c_str(), ss.str().c_str());
|
||||
// LOG_DEBUG("split prompt \"%s\" to tokens %s", original_text.c_str(), ss.str().c_str());
|
||||
printf("split prompt \"%s\" to tokens %s \n", original_text.c_str(), ss.str().c_str());
|
||||
return bpe_tokens;
|
||||
}
|
||||
};
|
||||
@@ -1255,10 +1286,10 @@ struct FrozenCLIPEmbedderWithCustomWords : public GGMLModule {
|
||||
}
|
||||
}
|
||||
|
||||
// for (int i = 0; i < tokens.size(); i++) {
|
||||
// std::cout << tokens[i] << ":" << weights[i] << ", ";
|
||||
// }
|
||||
// std::cout << std::endl;
|
||||
for (int i = 0; i < tokens.size(); i++) {
|
||||
std::cout << tokens[i] << ":" << weights[i] << ", ";
|
||||
}
|
||||
std::cout << std::endl;
|
||||
|
||||
return {tokens, weights};
|
||||
}
|
||||
|
||||
@@ -268,7 +268,7 @@ public:
|
||||
|
||||
// load weights
|
||||
LOG_DEBUG("loading weights");
|
||||
fprintf(stderr, "%s: loading weights \n", __func__);
|
||||
|
||||
int64_t t0 = ggml_time_ms();
|
||||
|
||||
std::map<std::string, struct ggml_tensor*> tensors_need_to_load;
|
||||
@@ -303,7 +303,7 @@ public:
|
||||
ggml_free(ctx);
|
||||
return false;
|
||||
}
|
||||
fprintf(stderr, "%s: loading weights sccess\n", __func__);
|
||||
|
||||
|
||||
// LOG_DEBUG("model size = %.2fMB", total_size / 1024.0 / 1024.0);
|
||||
|
||||
@@ -474,6 +474,110 @@ public:
|
||||
curr_lora_state = lora_state;
|
||||
}
|
||||
|
||||
|
||||
std::tuple<ggml_tensor*, ggml_tensor*, ggml_tensor*>
|
||||
get_learned_condition_with_trigger(ggml_context* work_ctx,
|
||||
const std::string& text,
|
||||
int clip_skip,
|
||||
int width,
|
||||
int height,
|
||||
int num_input_imgs,
|
||||
bool force_zero_embeddings = false) {
|
||||
cond_stage_model.set_clip_skip(clip_skip);
|
||||
auto image_token_id = cond_stage_model.tokenize(trigger_word, true);
|
||||
|
||||
|
||||
auto tokens_and_weights = cond_stage_model.tokenize(text, true);
|
||||
std::vector<int>& tokens = tokens_and_weights.first;
|
||||
std::vector<float>& weights = tokens_and_weights.second;
|
||||
int64_t t0 = ggml_time_ms();
|
||||
struct ggml_tensor* pooled = NULL;
|
||||
size_t total_hidden_size = cond_stage_model.text_model.hidden_size;
|
||||
if (version == VERSION_XL) {
|
||||
total_hidden_size += cond_stage_model.text_model2.hidden_size;
|
||||
pooled = ggml_new_tensor_1d(work_ctx, GGML_TYPE_F32, cond_stage_model.text_model2.projection_dim);
|
||||
}
|
||||
struct ggml_tensor* hidden_states = ggml_new_tensor_2d(work_ctx,
|
||||
GGML_TYPE_F32,
|
||||
total_hidden_size,
|
||||
cond_stage_model.text_model.max_position_embeddings); // [N, n_token, hidden_size]
|
||||
cond_stage_model.alloc_compute_buffer(work_ctx, (int)tokens.size());
|
||||
cond_stage_model.compute(n_threads, tokens, hidden_states, pooled);
|
||||
cond_stage_model.free_compute_buffer();
|
||||
// if (pooled != NULL) {
|
||||
// print_ggml_tensor(hidden_states);
|
||||
// print_ggml_tensor(pooled);
|
||||
// }
|
||||
|
||||
ggml_tensor *class_tokens_mask = ggml_dup_tensor(work_ctx, hidden_states);
|
||||
|
||||
int64_t t1 = ggml_time_ms();
|
||||
LOG_DEBUG("computing condition graph completed, taking %" PRId64 " ms", t1 - t0);
|
||||
ggml_tensor* result = ggml_dup_tensor(work_ctx, hidden_states);
|
||||
{
|
||||
float original_mean = ggml_tensor_mean(hidden_states);
|
||||
for (int i2 = 0; i2 < hidden_states->ne[2]; i2++) {
|
||||
for (int i1 = 0; i1 < hidden_states->ne[1]; i1++) {
|
||||
for (int i0 = 0; i0 < hidden_states->ne[0]; i0++) {
|
||||
float value = ggml_tensor_get_f32(hidden_states, i0, i1, i2);
|
||||
value *= weights[i1];
|
||||
ggml_tensor_set_f32(result, value, i0, i1, i2);
|
||||
}
|
||||
}
|
||||
}
|
||||
float new_mean = ggml_tensor_mean(result);
|
||||
ggml_tensor_scale(result, (original_mean / new_mean));
|
||||
}
|
||||
if (force_zero_embeddings) {
|
||||
float* vec = (float*)result->data;
|
||||
for (int i = 0; i < ggml_nelements(result); i++) {
|
||||
vec[i] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
ggml_tensor* vec = NULL;
|
||||
if (version == VERSION_XL) {
|
||||
int out_dim = 256;
|
||||
vec = ggml_new_tensor_1d(work_ctx, GGML_TYPE_F32, diffusion_model.adm_in_channels);
|
||||
// [0:1280]
|
||||
size_t offset = 0;
|
||||
memcpy(vec->data, pooled->data, ggml_nbytes(pooled));
|
||||
offset += ggml_nbytes(pooled);
|
||||
|
||||
struct ggml_tensor* timesteps = ggml_new_tensor_1d(work_ctx, GGML_TYPE_F32, 2);
|
||||
// original_size_as_tuple
|
||||
float orig_width = (float)width;
|
||||
float orig_height = (float)height;
|
||||
ggml_tensor_set_f32(timesteps, orig_height, 0);
|
||||
ggml_tensor_set_f32(timesteps, orig_width, 1);
|
||||
ggml_tensor* embed_view = ggml_view_2d(work_ctx, vec, out_dim, 2, ggml_type_size(GGML_TYPE_F32) * out_dim, offset);
|
||||
offset += ggml_nbytes(embed_view);
|
||||
set_timestep_embedding(timesteps, embed_view, out_dim);
|
||||
// print_ggml_tensor(ggml_reshape_1d(work_ctx, embed_view, out_dim * 2));
|
||||
// crop_coords_top_left
|
||||
float crop_coord_top = 0.f;
|
||||
float crop_coord_left = 0.f;
|
||||
ggml_tensor_set_f32(timesteps, crop_coord_top, 0);
|
||||
ggml_tensor_set_f32(timesteps, crop_coord_left, 1);
|
||||
embed_view = ggml_view_2d(work_ctx, vec, out_dim, 2, ggml_type_size(GGML_TYPE_F32) * out_dim, offset);
|
||||
offset += ggml_nbytes(embed_view);
|
||||
set_timestep_embedding(timesteps, embed_view, out_dim);
|
||||
// print_ggml_tensor(ggml_reshape_1d(work_ctx, embed_view, out_dim * 2));
|
||||
// target_size_as_tuple
|
||||
float target_width = (float)width;
|
||||
float target_height = (float)height;
|
||||
ggml_tensor_set_f32(timesteps, target_height, 0);
|
||||
ggml_tensor_set_f32(timesteps, target_width, 1);
|
||||
embed_view = ggml_view_2d(work_ctx, vec, out_dim, 2, ggml_type_size(GGML_TYPE_F32) * out_dim, offset);
|
||||
offset += ggml_nbytes(embed_view);
|
||||
set_timestep_embedding(timesteps, embed_view, out_dim);
|
||||
// print_ggml_tensor(ggml_reshape_1d(work_ctx, embed_view, out_dim * 2));
|
||||
GGML_ASSERT(offset == ggml_nbytes(vec));
|
||||
}
|
||||
// print_ggml_tensor(result);
|
||||
return std::make_tuple(result, vec, class_tokens_mask);
|
||||
}
|
||||
|
||||
std::pair<ggml_tensor*, ggml_tensor*> get_learned_condition(ggml_context* work_ctx,
|
||||
const std::string& text,
|
||||
int clip_skip,
|
||||
|
||||
Reference in New Issue
Block a user