From 5be4787a12c42a71e00f709ee701ea505464c7ea Mon Sep 17 00:00:00 2001 From: Michael Poutre Date: Mon, 28 Aug 2023 22:01:16 -0700 Subject: [PATCH] fix(ClipTransformer): Update with transformers library --- lib/fun_clip_stuff.py | 16 ++++++++++------ 1 file changed, 10 insertions(+), 6 deletions(-) diff --git a/lib/fun_clip_stuff.py b/lib/fun_clip_stuff.py index 8730eaa..fcbfaba 100644 --- a/lib/fun_clip_stuff.py +++ b/lib/fun_clip_stuff.py @@ -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]