fix(ClipTransformer): Update with transformers library

This commit is contained in:
Michael Poutre
2023-08-28 22:01:16 -07:00
parent fd2316c8fa
commit 5be4787a12
+10 -6
View File
@@ -3,8 +3,13 @@ from typing import Optional, Tuple, Union
import torch
from transformers import CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from transformers.models.clip.modeling_clip import (
_expand_mask,
CLIPTextEmbeddings,
CLIPTextTransformer,
CLIPTextModel,
_make_causal_mask,
)
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
@@ -97,12 +102,11 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
bsz = len(input_ids)
# TODO: Properly gather this
seq_len = 77
# bsz, seq_len = input_shape
input_shape = torch.Size([bsz, seq_len])
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device
)
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
# expand attention_mask
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]