From 30ad69fb9102e9ba1effc6b0f4b4ef119c10ed2f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 17 Mar 2024 20:13:59 +0200 Subject: [PATCH] Disable dual patchnorm to support latest open_clip_torch --- lvdm/modules/encoders/condition.py | 20 ++++++++++---------- requirements.txt | 2 +- 2 files changed, 11 insertions(+), 11 deletions(-) diff --git a/lvdm/modules/encoders/condition.py b/lvdm/modules/encoders/condition.py index e70bbe2..2bf0a2b 100644 --- a/lvdm/modules/encoders/condition.py +++ b/lvdm/modules/encoders/condition.py @@ -343,17 +343,17 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder): x = self.preprocess(x) # to patches - whether to use dual patchnorm - https://arxiv.org/abs/2302.01327v1 - if self.model.visual.input_patchnorm: + #if self.model.visual.input_patchnorm: # einops - rearrange(x, 'b c (h p1) (w p2) -> b (h w) (c p1 p2)') - x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1]) - x = x.permute(0, 2, 4, 1, 3, 5) - x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1) - x = self.model.visual.patchnorm_pre_ln(x) - x = self.model.visual.conv1(x) - else: - x = self.model.visual.conv1(x) # shape = [*, width, grid, grid] - x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2] - x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] + # x = x.reshape(x.shape[0], x.shape[1], self.model.visual.grid_size[0], self.model.visual.patch_size[0], self.model.visual.grid_size[1], self.model.visual.patch_size[1]) + # x = x.permute(0, 2, 4, 1, 3, 5) + # x = x.reshape(x.shape[0], self.model.visual.grid_size[0] * self.model.visual.grid_size[1], -1) + # x = self.model.visual.patchnorm_pre_ln(x) + # x = self.model.visual.conv1(x) + #else: + x = self.model.visual.conv1(x) # shape = [*, width, grid, grid] + x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2] + x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width] # class embeddings and positional embeddings x = torch.cat( diff --git a/requirements.txt b/requirements.txt index 7a7f76c..f05104d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,5 +6,5 @@ pytorch_lightning>=1.8.3 tqdm>=4.65.0 transformers>=4.25.1 timm -open_clip_torch==2.12.0 +open_clip_torch>=2.12.0 kornia \ No newline at end of file