From 304704fa36c80ef2c0413b4c82c497f2292f6371 Mon Sep 17 00:00:00 2001 From: bssrdf Date: Thu, 8 Feb 2024 22:16:01 -0500 Subject: [PATCH] corrected clip vision model output --- clip.hpp | 18 +++++++++++------ pmid.hpp | 59 +++++++++++++++++++++++++++++++++++++++++++++++++++++--- 2 files changed, 68 insertions(+), 9 deletions(-) diff --git a/clip.hpp b/clip.hpp index 769deb94..b2fb0264 100644 --- a/clip.hpp +++ b/clip.hpp @@ -945,7 +945,9 @@ struct CLIPVisionModel { return h; } - struct ggml_tensor* forward(struct ggml_context* ctx0, struct ggml_tensor *x, + struct ggml_tensor* forward(struct ggml_context* ctx0, + struct ggml_tensor *x, + struct ggml_tensor *cls, struct ggml_tensor *temp, struct ggml_tensor *positions) { // input_ids: [N, n_token] @@ -978,8 +980,6 @@ struct CLIPVisionModel { // struct ggml_tensor * positions = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, num_positions); - - embeddings = ggml_add(ctx0, embeddings, ggml_repeat(ctx0, ggml_get_rows(ctx0, position_embeddings, positions), embeddings)); @@ -991,12 +991,18 @@ struct CLIPVisionModel { embeddings = resblocks[i].forward(ctx0, embeddings); // [N, n_token, hidden_size] } + // get the output of cls token, e.g., 0th index + embeddings = ggml_get_rows(ctx0, ggml_reshape_2d(ctx0, embeddings, hidden_size, num_positions * batch_size), cls); + + // post-layernorm + embeddings = ggml_nn_layer_norm(ctx0, embeddings, post_ln_w, post_ln_b); + struct ggml_tensor * cur = embeddings; - ggml_set_name(cur, "last_hidden_state"); - // we return last hidden state instead of image embedding here + ggml_set_name(cur, "pooled_output"); + // we return pooled output instead of image embedding here // TODO: add other return items - return cur; // [batch_size, seq_length, hidden_size] + return cur; // [batch_size, hidden_size] } diff --git a/pmid.hpp b/pmid.hpp index 64ac872a..57e1d0eb 100644 --- a/pmid.hpp +++ b/pmid.hpp @@ -173,6 +173,36 @@ struct FuseModule{ // x: [N, channels, h, w] // in_layers + + // # id_embeds shape: [b, max_num_inputs, 1, 2048] + // id_embeds = id_embeds.to(prompt_embeds.dtype) + // num_inputs = class_tokens_mask.sum().unsqueeze(0) # TODO: check for training case + // batch_size, max_num_inputs = id_embeds.shape[:2] + // # seq_length: 77 + // seq_length = prompt_embeds.shape[1] + // # flat_id_embeds shape: [b*max_num_inputs, 1, 2048] + // flat_id_embeds = id_embeds.view( + // -1, id_embeds.shape[-2], id_embeds.shape[-1] + // ) + // # valid_id_mask [b*max_num_inputs] + // valid_id_mask = ( + // torch.arange(max_num_inputs, device=flat_id_embeds.device)[None, :] + // < num_inputs[:, None] + // ) + // valid_id_embeds = flat_id_embeds[valid_id_mask.flatten()] + + // 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]) + // # 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) + // 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; } @@ -235,6 +265,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* cls, struct ggml_tensor* class_embedding_temp, struct ggml_tensor* positions) { // x: [N, channels, h, w] @@ -244,6 +275,7 @@ struct PhotoMakerIDEncoder : public GGMLModule { print_ggml_tensor(class_embedding_temp, true, "class_embedding_temp"); struct ggml_tensor *shared_id_embeds = vision_model.forward(ctx, id_pixel_values, + cls, class_embedding_temp, positions ); // [batch_size, seq_length, hidden_size] @@ -256,7 +288,20 @@ struct PhotoMakerIDEncoder : public GGMLModule { // id_embeds = torch.cat((id_embeds, id_embeds_2), dim=-1) print_ggml_tensor(id_embeds, true, "id_embeds"); print_ggml_tensor(id_embeds_2, true, "id_embeds_2"); + + id_embeds = ggml_cont(ctx, ggml_permute(ctx, id_embeds, 2, 0, 1, 3)); + id_embeds_2 = ggml_cont(ctx, ggml_permute(ctx, id_embeds_2, 2, 0, 1, 3)); + print_ggml_tensor(id_embeds, true, "id_embeds_after_perm"); + print_ggml_tensor(id_embeds_2, true, "id_embeds_2_after_perm"); + id_embeds = ggml_concat(ctx, id_embeds, id_embeds_2); // [batch_size, seq_length, 1, 2048] check whether concat at dim 2 is right + print_ggml_tensor(id_embeds, true, "id_embeds_after_cat"); + + + id_embeds = ggml_cont(ctx, ggml_permute(ctx, id_embeds, 1, 2, 0, 3)); + 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); return updated_prompt_embeds; @@ -294,10 +339,12 @@ struct PhotoMakerIDEncoder : public GGMLModule { const int num_patches = ((image_size / vision_model.patch_size) * (image_size / vision_model.patch_size)); const int num_positions = num_patches + 1; - + struct ggml_tensor * cls = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, batch_size); + struct ggml_tensor * class_embedding_temp = ggml_new_tensor_4d(ctx0, GGML_TYPE_F32, vision_model.hidden_size, 1, 1, batch_size); struct ggml_tensor * positions = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, num_positions); + ggml_allocr_alloc(allocr, cls); ggml_allocr_alloc(allocr, class_embedding_temp); ggml_allocr_alloc(allocr, positions); @@ -307,16 +354,22 @@ struct PhotoMakerIDEncoder : public GGMLModule { 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)); + std::vector cls_h; + for (int b = 0; b < batch_size; b++) { + cls_h.push_back(b * num_positions); + } std::vector pos; for (int i = 0; i < num_positions; i++) { pos.push_back(i); } + ggml_backend_tensor_set(cls, cls_h.data(), 0, ggml_nbytes(cls)); ggml_backend_tensor_set(positions, pos.data(), 0, ggml_nbytes(positions)); } struct ggml_tensor* updated_prompt_embeds = forward(ctx0, id_pixel_values_d, - prompt_embeds_d, - class_tokens_mask_d, + prompt_embeds_d, + class_tokens_mask_d, + cls, class_embedding_temp, positions );