This commit is contained in:
parent
1be09952fc
commit
22749c8dc1
|
|
@ -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()
|
||||||
|
|
|
||||||
|
|
@ -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)
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue