diff --git a/train/hash_train.py b/train/hash_train.py index 7f5b82e..6ed4d16 100644 --- a/train/hash_train.py +++ b/train/hash_train.py @@ -14,7 +14,7 @@ from torch.nn import functional as F from utils import get_args, calc_neighbor, cosine_similarity, euclidean_similarity,find_indices from utils.calc_utils import cal_map, cal_pr from dataset.dataloader import dataloader -import open_clip +import clip # from transformers import BertModel device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") @@ -41,7 +41,7 @@ class Trainer(TrainBase): def _init_model(self): self.logger.info("init model.") - model_clip, _, preprocess = open_clip.create_model_and_transforms('ViT-B-16', device=device) + model_clip, preprocess = clip.load(self.args.victim, device=device) self.model= model_clip self.model.eval() self.model.float() diff --git a/utils/get_args.py b/utils/get_args.py index 42b6a22..bb23590 100644 --- a/utils/get_args.py +++ b/utils/get_args.py @@ -15,7 +15,7 @@ def get_args(): parser.add_argument("--label-file", type=str, default="label.mat") parser.add_argument("--similarity-function", type=str, default="euclidean", help="choise form [cosine, euclidean]") parser.add_argument("--loss-type", type=str, default="l2", help="choise form [l1, l2]") - # parser.add_argument("--test-index-file", type=str, default="./data/test/index.mat") + parser.add_argument('--victim', default='ViT-B/16', choices=['ViT-L/14', 'ViT-B/16', 'ViT-B/32', 'RN50', 'RN101']) # parser.add_argument("--test-caption-file", type=str, default="./data/test/captions.mat") # parser.add_argument("--test-label-file", type=str, default="./data/test/label.mat") parser.add_argument("--txt-dim", type=int, default=1024)