continue making progress in id fusion process

This commit is contained in:
bssrdf
2024-02-09 16:55:29 -05:00
parent 304704fa36
commit 184e4b84a2
2 changed files with 69 additions and 34 deletions

View File

@@ -79,7 +79,7 @@ struct FuseBlock {
// x: [N, channels, h, w]
// in_layers
auto h = ggml_nn_group_norm(ctx, x, ln_w, ln_b);
auto h = ggml_nn_layer_norm(ctx, x, ln_w, ln_b);
h = ggml_add(ctx, ggml_mul_mat(ctx, fc1_w, h), fc1_b);
h = ggml_gelu_inplace(ctx, h);
h = ggml_add(ctx, ggml_mul_mat(ctx, fc2_w, h), fc2_b);
@@ -105,11 +105,6 @@ struct FuseModule{
embed_dim(imb_d),
mlp1(imb_d*2, imb_d, imb_d, false),
mlp2(imb_d, imb_d, imb_d, true) {
// mlp1 = FuseBlock(embed_dim*2, embed_dim, embed_dim, false);
// mlp2 = FuseBlock(embed_dim*2, embed_dim, embed_dim, true);
}
@@ -154,13 +149,23 @@ struct FuseModule{
struct ggml_tensor* id_embeds) {
// x: [N, channels, h, w]
// stacked_id_embeds = torch.cat([prompt_embeds, id_embeds], dim=-1)
// stacked_id_embeds = self.mlp1(stacked_id_embeds) + prompt_embeds
// stacked_id_embeds = self.mlp2(stacked_id_embeds)
// stacked_id_embeds = self.layer_norm(stacked_id_embeds)
// return stacked_id_embeds
// in_layers
auto stacked_id_embeds = ggml_concat(ctx, prompt_embeds, id_embeds); // check whether concat at dim 2 is right
auto prompt_embeds0 = ggml_cont(ctx, ggml_permute(ctx, prompt_embeds, 2, 0, 1, 3));
auto id_embeds0 = ggml_cont(ctx, ggml_permute(ctx, id_embeds, 2, 0, 1, 3));
// concat is along dim 2
auto stacked_id_embeds = ggml_concat(ctx, prompt_embeds0, id_embeds0);
stacked_id_embeds = ggml_cont(ctx, ggml_permute(ctx, stacked_id_embeds, 1, 2, 0, 3));
stacked_id_embeds = mlp1.forward(ctx, stacked_id_embeds);
stacked_id_embeds = ggml_add(ctx, stacked_id_embeds, prompt_embeds);
stacked_id_embeds = mlp2.forward(ctx, stacked_id_embeds);
stacked_id_embeds = ggml_nn_group_norm(ctx, stacked_id_embeds, ln_w, ln_b);
stacked_id_embeds = ggml_nn_layer_norm(ctx, stacked_id_embeds, ln_w, ln_b);
return stacked_id_embeds;
}
@@ -169,7 +174,8 @@ struct FuseModule{
struct ggml_tensor* forward(struct ggml_context* ctx,
struct ggml_tensor* prompt_embeds,
struct ggml_tensor* id_embeds,
struct ggml_tensor* class_tokens_mask) {
struct ggml_tensor* class_tokens_mask,
struct ggml_tensor* class_tokens_mask_pos) {
// x: [N, channels, h, w]
// in_layers
@@ -193,16 +199,16 @@ struct FuseModule{
// prompt_embeds = prompt_embeds.view(-1, prompt_embeds.shape[-1])
// class_tokens_mask = class_tokens_mask.view(-1)
// valid_id_embeds = valid_id_embeds.view(-1, valid_id_embeds.shape[-1])
struct ggml_tensor * valid_id_embeds = id_embeds;
// # slice out the image token embeddings
// image_token_embeds = prompt_embeds[class_tokens_mask]
// stacked_id_embeds = self.fuse_fn(image_token_embeds, valid_id_embeds)
struct ggml_tensor * image_token_embeds = ggml_get_rows(ctx, prompt_embeds, class_tokens_mask_pos);
print_ggml_tensor(image_token_embeds, true, "image_token_embeds");
struct ggml_tensor *stacked_id_embeds = fuse_fn(ctx, image_token_embeds, valid_id_embeds);
print_ggml_tensor(stacked_id_embeds, true, "stacked_id_embeds");
// assert class_tokens_mask.sum() == stacked_id_embeds.shape[0], f"{class_tokens_mask.sum()} != {stacked_id_embeds.shape[0]}"
// prompt_embeds.masked_scatter_(class_tokens_mask[:, None], stacked_id_embeds.to(prompt_embeds.dtype))
// updated_prompt_embeds = prompt_embeds.view(batch_size, seq_length, -1)
// return updated_prompt_embeds
int64_t *ne = id_embeds->ne;
// struct ggml_tensor* flat_id_embeds = ggml_view_3d(ctx, id_embeds, ne[0], 1, ne[1]*ne[2], );
struct ggml_tensor* h = NULL;
return h;
}
@@ -265,6 +271,7 @@ struct PhotoMakerIDEncoder : public GGMLModule {
struct ggml_tensor* id_pixel_values,
struct ggml_tensor* prompt_embeds,
struct ggml_tensor* class_tokens_mask,
struct ggml_tensor* class_tokens_mask_pos,
struct ggml_tensor* cls,
struct ggml_tensor* class_embedding_temp,
struct ggml_tensor* positions) {
@@ -280,6 +287,16 @@ struct PhotoMakerIDEncoder : public GGMLModule {
positions
); // [batch_size, seq_length, hidden_size]
print_ggml_tensor(shared_id_embeds, true, "shared_id_embeds");
if(class_tokens_mask->backend == GGML_BACKEND_GPU){
int *ctm = (int *)malloc(class_tokens_mask->ne[0]);
ggml_backend_tensor_get(class_tokens_mask, ctm, 0, ggml_nbytes(class_tokens_mask));
printf("class_tokens_mask[");
for(int i = 0; i < class_tokens_mask->ne[0]; i++)
printf("%d, ", ctm[i]);
printf("]\n");
free(ctm);
}
struct ggml_tensor *id_embeds = vision_model.visual_project(ctx, shared_id_embeds); // [batch_size, seq_length, proj_dim(768)]
struct ggml_tensor *id_embeds_2 = ggml_mul_mat(ctx, visual_projection_2, shared_id_embeds); // [batch_size, seq_length, 1280]
// id_embeds = id_embeds.view(b, num_inputs, 1, -1)
@@ -302,7 +319,8 @@ struct PhotoMakerIDEncoder : public GGMLModule {
print_ggml_tensor(id_embeds, true, "id_embeds_after_cat+perm");
struct ggml_tensor * updated_prompt_embeds = fuse_module.forward(ctx, prompt_embeds, id_embeds, class_tokens_mask);
struct ggml_tensor * updated_prompt_embeds = fuse_module.forward(ctx, prompt_embeds, id_embeds,
class_tokens_mask, class_tokens_mask_pos);
return updated_prompt_embeds;
@@ -311,7 +329,7 @@ struct PhotoMakerIDEncoder : public GGMLModule {
struct ggml_cgraph* build_graph(struct ggml_allocr* allocr,
struct ggml_tensor* id_pixel_values,
struct ggml_tensor* prompt_embeds,
struct ggml_tensor* class_tokens_mask
std::vector<bool> &class_tokens_mask
) {
// since we are using ggml-alloc, this buffer only needs enough space to hold the ggml_tensor and ggml_cgraph structs, but not the tensor data
static size_t buf_size = ggml_tensor_overhead() * GGML_DEFAULT_GRAPH_SIZE + ggml_graph_overhead();
@@ -331,9 +349,23 @@ struct PhotoMakerIDEncoder : public GGMLModule {
ggml_allocr_alloc(allocr, id_pixel_values_d);
struct ggml_tensor* prompt_embeds_d = ggml_dup_tensor(ctx0, prompt_embeds);
ggml_allocr_alloc(allocr, prompt_embeds_d);
struct ggml_tensor* class_tokens_mask_d = ggml_dup_tensor(ctx0, class_tokens_mask);
struct ggml_tensor* class_tokens_mask_d = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, class_tokens_mask.size());
ggml_allocr_alloc(allocr, class_tokens_mask_d);
std::vector<int> ctm;
std::vector<int> ctmpos;
for(int i=0; i < class_tokens_mask.size(); i++){
if(class_tokens_mask[i]){
ctm.push_back(1);
ctmpos.push_back(i);
}else{
ctm.push_back(0);
}
}
struct ggml_tensor* class_tokens_mask_pos = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, ctmpos.size());
ggml_allocr_alloc(allocr, class_tokens_mask_pos);
const int image_size = id_pixel_values->ne[0];
int batch_size = id_pixel_values->ne[3];
const int num_patches = ((image_size / vision_model.patch_size) * (image_size / vision_model.patch_size));
@@ -353,7 +385,7 @@ struct PhotoMakerIDEncoder : public GGMLModule {
if (!ggml_allocr_is_measure(allocr)) {
ggml_backend_tensor_set(id_pixel_values_d, id_pixel_values->data, 0, ggml_nbytes(id_pixel_values));
ggml_backend_tensor_set(prompt_embeds_d, prompt_embeds->data, 0, ggml_nbytes(prompt_embeds));
ggml_backend_tensor_set(class_tokens_mask_d, class_tokens_mask->data, 0, ggml_nbytes(class_tokens_mask_d));
ggml_backend_tensor_set(class_tokens_mask_d, ctm.data(), 0, ggml_nbytes(class_tokens_mask_d));
std::vector<int> cls_h;
for (int b = 0; b < batch_size; b++) {
cls_h.push_back(b * num_positions);
@@ -364,11 +396,13 @@ struct PhotoMakerIDEncoder : public GGMLModule {
}
ggml_backend_tensor_set(cls, cls_h.data(), 0, ggml_nbytes(cls));
ggml_backend_tensor_set(positions, pos.data(), 0, ggml_nbytes(positions));
ggml_backend_tensor_set(class_tokens_mask_pos, ctmpos.data(), 0, ggml_nbytes(class_tokens_mask_pos));
}
struct ggml_tensor* updated_prompt_embeds = forward(ctx0,
id_pixel_values_d,
prompt_embeds_d,
class_tokens_mask_d,
class_tokens_mask_pos,
cls,
class_embedding_temp,
positions
@@ -383,7 +417,7 @@ struct PhotoMakerIDEncoder : public GGMLModule {
void alloc_compute_buffer(ggml_context* work_ctx,
struct ggml_tensor* id_pixel_values,
struct ggml_tensor* prompt_embeds,
struct ggml_tensor* class_tokens_mask) {
std::vector<bool> &class_tokens_mask) {
auto get_graph = [&]() -> struct ggml_cgraph* {
return build_graph(compute_allocr, id_pixel_values, prompt_embeds, class_tokens_mask);
@@ -394,7 +428,7 @@ struct PhotoMakerIDEncoder : public GGMLModule {
void compute(const int n_threads,
struct ggml_tensor* id_pixel_values,
struct ggml_tensor* prompt_embeds,
struct ggml_tensor* class_tokens_mask,
std::vector<bool> &class_tokens_mask,
struct ggml_tensor* updated_prompt_embeds) {
auto get_graph = [&]() -> struct ggml_cgraph* {

View File

@@ -489,7 +489,7 @@ public:
}
std::tuple<ggml_tensor*, ggml_tensor*, ggml_tensor*>
std::tuple<ggml_tensor*, ggml_tensor*, std::vector<bool>>
get_learned_condition_with_trigger(ggml_context* work_ctx,
const std::string& text,
int clip_skip,
@@ -531,13 +531,13 @@ public:
// print_ggml_tensor(pooled);
// }
ggml_tensor *class_tokens_mask = ggml_new_tensor_1d(work_ctx, GGML_TYPE_I32, cond_stage_model.text_model.max_position_embeddings);
for (int i1 = 0; i1 < hidden_states->ne[1]; i1++) {
if(clsm[i1])
ggml_set_i32_1d(class_tokens_mask, i1, 1);
else
ggml_set_i32_1d(class_tokens_mask, i1, 0);
}
// ggml_tensor *class_tokens_mask = ggml_new_tensor_1d(work_ctx, GGML_TYPE_I32, cond_stage_model.text_model.max_position_embeddings);
// for (int i1 = 0; i1 < hidden_states->ne[1]; i1++) {
// if(clsm[i1])
// ggml_set_i32_1d(class_tokens_mask, i1, 1);
// else
// ggml_set_i32_1d(class_tokens_mask, i1, 0);
// }
int64_t t1 = ggml_time_ms();
LOG_DEBUG("computing condition graph completed, taking %" PRId64 " ms", t1 - t0);
@@ -603,13 +603,13 @@ public:
GGML_ASSERT(offset == ggml_nbytes(vec));
}
// print_ggml_tensor(result);
return std::make_tuple(result, vec, class_tokens_mask);
return std::make_tuple(result, vec, clsm);
}
ggml_tensor* id_encoder(ggml_context* work_ctx,
ggml_tensor* init_img,
ggml_tensor* prompts_embeds,
ggml_tensor* class_tokens_mask){
std::vector<bool> &class_tokens_mask){
size_t total_hidden_size = cond_stage_model.text_model.hidden_size;
if (version == VERSION_XL) {
@@ -1455,7 +1455,8 @@ sd_image_t* txt2img(sd_ctx_t* sd_ctx,
ggml_tensor* init_img = NULL;
ggml_tensor* prompts_embeds = NULL;
ggml_tensor* pooled_prompts_embeds = NULL;
ggml_tensor* class_tokens_mask = NULL;
// ggml_tensor* class_tokens_mask = NULL;
std::vector<bool> class_tokens_mask;
if(sd_ctx->sd->stacked_id){
int32_t width = input_id_images[0]->width;
int32_t height = input_id_images[0]->height;
@@ -1483,9 +1484,9 @@ sd_image_t* txt2img(sd_ctx_t* sd_ctx,
ne = pooled_prompts_embeds->ne;
fprintf(stderr, "%s: SDXL pooled text embedding tensor ne [%ld, %ld, %ld, %ld] \n",
__func__, ne[0], ne[1], ne[2], ne[3]);
ne = class_tokens_mask->ne;
fprintf(stderr, "%s: class token mask tensor ne [%ld, %ld, %ld, %ld] \n",
__func__, ne[0], ne[1], ne[2], ne[3]);
// ne = class_tokens_mask->ne;
// fprintf(stderr, "%s: class token mask tensor ne [%ld, %ld, %ld, %ld] \n",
// __func__, ne[0], ne[1], ne[2], ne[3]);
prompts_embeds = sd_ctx->sd->id_encoder(work_ctx, init_img, prompts_embeds, class_tokens_mask);