This commit is contained in:
leewlving 2024-06-17 20:59:34 +08:00
parent 1be09952fc
commit 22749c8dc1
2 changed files with 3 additions and 3 deletions

View File

@ -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 import get_args, calc_neighbor, cosine_similarity, euclidean_similarity,find_indices
from utils.calc_utils import cal_map, cal_pr from utils.calc_utils import cal_map, cal_pr
from dataset.dataloader import dataloader from dataset.dataloader import dataloader
import open_clip import clip
# from transformers import BertModel # from transformers import BertModel
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@ -41,7 +41,7 @@ class Trainer(TrainBase):
def _init_model(self): def _init_model(self):
self.logger.info("init model.") 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= model_clip
self.model.eval() self.model.eval()
self.model.float() self.model.float()

View File

@ -15,7 +15,7 @@ def get_args():
parser.add_argument("--label-file", type=str, default="label.mat") 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("--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("--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-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("--test-label-file", type=str, default="./data/test/label.mat")
parser.add_argument("--txt-dim", type=int, default=1024) parser.add_argument("--txt-dim", type=int, default=1024)