mirror of
https://github.com/leejet/stable-diffusion.cpp.git
synced 2026-07-29 22:30:43 -05:00
corrected clip vision model output
This commit is contained in:
18
clip.hpp
18
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]
|
||||
|
||||
}
|
||||
|
||||
|
||||
59
pmid.hpp
59
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<int> cls_h;
|
||||
for (int b = 0; b < batch_size; b++) {
|
||||
cls_h.push_back(b * num_positions);
|
||||
}
|
||||
std::vector<int> 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
|
||||
);
|
||||
|
||||
Reference in New Issue
Block a user