use our blip

This commit is contained in:
lllyasviel
2023-12-12 21:07:39 -08:00
parent c175afb394
commit 322aa5a724
10 changed files with 25 additions and 31 deletions
View File
+2 -2
View File
@@ -8,8 +8,8 @@
import warnings
warnings.filterwarnings("ignore")
from models.vit import VisionTransformer, interpolate_pos_embed
from models.med import BertConfig, BertModel, BertLMHeadModel
from extras.BLIP.models.vit import VisionTransformer, interpolate_pos_embed
from extras.BLIP.models.med import BertConfig, BertModel, BertLMHeadModel
from transformers import BertTokenizer
import torch
+2 -2
View File
@@ -1,11 +1,11 @@
from models.med import BertConfig, BertModel
from extras.BLIP.models.med import BertConfig, BertModel
from transformers import BertTokenizer
import torch
from torch import nn
import torch.nn.functional as F
from models.blip import create_vit, init_tokenizer, load_checkpoint
from extras.BLIP.models.blip import create_vit, init_tokenizer, load_checkpoint
class BLIP_ITM(nn.Module):
def __init__(self,
+6 -4
View File
@@ -1,7 +1,7 @@
from models.med import BertConfig
from models.nlvr_encoder import BertModel
from models.vit import interpolate_pos_embed
from models.blip import create_vit, init_tokenizer, is_url
from extras.BLIP.models.med import BertConfig
from extras.BLIP.models.nlvr_encoder import BertModel
from extras.BLIP.models.vit import interpolate_pos_embed
from extras.BLIP.models.blip import create_vit, init_tokenizer, is_url
from timm.models.hub import download_cached_file
@@ -10,6 +10,8 @@ from torch import nn
import torch.nn.functional as F
from transformers import BertTokenizer
import numpy as np
import os
class BLIP_NLVR(nn.Module):
def __init__(self,
+3 -3
View File
@@ -5,7 +5,7 @@
* For full license text, see LICENSE.txt file in the repo root or https://opensource.org/licenses/BSD-3-Clause
* By Junnan Li
'''
from models.med import BertConfig, BertModel, BertLMHeadModel
from extras.BLIP.models.med import BertConfig, BertModel, BertLMHeadModel
from transformers import BertTokenizer
import transformers
transformers.logging.set_verbosity_error()
@@ -14,7 +14,7 @@ import torch
from torch import nn
import torch.nn.functional as F
from models.blip import create_vit, init_tokenizer, load_checkpoint
from extras.BLIP.models.blip import create_vit, init_tokenizer, load_checkpoint
class BLIP_Pretrain(nn.Module):
def __init__(self,
@@ -270,7 +270,7 @@ from typing import List
def tie_encoder_decoder_weights(encoder: nn.Module, decoder: nn.Module, base_model_prefix: str, skip_key:str):
uninitialized_encoder_weights: List[str] = []
if decoder.__class__ != encoder.__class__:
logger.info(
print(
f"{decoder.__class__} and {encoder.__class__} are not equal. In this case make sure that all encoder weights are correctly initialized."
)
+2 -2
View File
@@ -1,11 +1,11 @@
from models.med import BertConfig, BertModel
from extras.BLIP.models.med import BertConfig, BertModel
from transformers import BertTokenizer
import torch
from torch import nn
import torch.nn.functional as F
from models.blip import create_vit, init_tokenizer, load_checkpoint
from extras.BLIP.models.blip import create_vit, init_tokenizer, load_checkpoint
class BLIP_Retrieval(nn.Module):
def __init__(self,
+2 -2
View File
@@ -1,5 +1,5 @@
from models.med import BertConfig, BertModel, BertLMHeadModel
from models.blip import create_vit, init_tokenizer, load_checkpoint
from extras.BLIP.models.med import BertConfig, BertModel, BertLMHeadModel
from extras.BLIP.models.blip import create_vit, init_tokenizer, load_checkpoint
import torch
from torch import nn
+4 -1
View File
@@ -18,7 +18,10 @@ from timm.models.registry import register_model
from timm.models.layers import trunc_normal_, DropPath
from timm.models.helpers import named_apply, adapt_input_conv
from fairscale.nn.checkpoint.checkpoint_activations import checkpoint_wrapper
def checkpoint_wrapper(x):
return x
class Mlp(nn.Module):
""" MLP as used in Vision Transformer, MLP-Mixer and related networks