From a83e26db8ab899b44b4d4014094110a00dd0554a Mon Sep 17 00:00:00 2001 From: hanjiang Date: Sun, 13 Jul 2025 09:00:09 +0800 Subject: [PATCH 1/9] =?UTF-8?q?=E2=9C=A8=20feat:=20init=20TPN?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config/TPN.yaml | 32 ++++ core/model/backbone/tpn_encoder.py | 49 +++++ core/model/metric/__init__.py | 3 +- core/model/metric/tpn.py | 174 ++++++++++++++++++ .../TPN-miniImageNet--ravi-10-1-Table1.yaml | 26 +++ .../TPN-miniImageNet--ravi-10-5-Table1.yaml | 26 +++ .../TPN-miniImageNet--ravi-5-1-Table1.yaml | 26 +++ .../TPN-miniImageNet--ravi-5-5-Table1.yaml | 26 +++ .../TPN-tieredImageNet--ravi-10-1-Table2.yaml | 26 +++ .../TPN-tieredImageNet--ravi-10-5-Table2.yaml | 26 +++ .../TPN-tieredImageNet--ravi-5-1-Table2.yaml | 26 +++ .../TPN-tieredImageNet--ravi-5-5-Table2.yaml | 26 +++ 12 files changed, 465 insertions(+), 1 deletion(-) create mode 100644 config/TPN.yaml create mode 100644 core/model/backbone/tpn_encoder.py create mode 100644 core/model/metric/tpn.py create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml diff --git a/config/TPN.yaml b/config/TPN.yaml new file mode 100644 index 00000000..63dfcba5 --- /dev/null +++ b/config/TPN.yaml @@ -0,0 +1,32 @@ +backbone: + name: CNNEncoder + kwargs: null + +classifier: + name: TPN + kwargs: + topk: 20 + sigma: 0.25 + alpha: 0.99 + rn: 300 + +way_num: 5 +shot_num: 1 +query_num: 15 + +epoch: 2100 +test_epoch: 100 +train_episode: 100 +test_episode: 100 +episode_size: 1 + +optimizer: + name: Adam + kwargs: + lr: 1e-3 + +lr_scheduler: + name: StepLR + kwargs: + step_size: 10000 + gamma: 0.5 \ No newline at end of file diff --git a/core/model/backbone/tpn_encoder.py b/core/model/backbone/tpn_encoder.py new file mode 100644 index 00000000..42f9c2b0 --- /dev/null +++ b/core/model/backbone/tpn_encoder.py @@ -0,0 +1,49 @@ +#------------------------------------- +# Project: Transductive Propagation Network for Few-shot Learning +# Date: 2019.1.11 +# Author: Yanbin Liu +# All Rights Reserved +#------------------------------------- + +import torch +import torch.nn as nn + +class CNNEncoder(nn.Module): + """Encoder for feature embedding""" + def __init__(self): + super(CNNEncoder, self).__init__() + self.layer1 = nn.Sequential( + nn.Conv2d(3, 64, kernel_size=3, padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(2)) + self.layer2 = nn.Sequential( + nn.Conv2d(64,64,kernel_size=3,padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(2)) + self.layer3 = nn.Sequential( + nn.Conv2d(64,64,kernel_size=3,padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(2)) + self.layer4 = nn.Sequential( + nn.Conv2d(64,64,kernel_size=3,padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(2)) + + def forward(self,x): + """x: bs*3*84*84 """ + out = self.layer1(x) + out = self.layer2(out) + out = self.layer3(out) + out = self.layer4(out) + + return out + + + + + + diff --git a/core/model/metric/__init__.py b/core/model/metric/__init__.py index 459d042e..1f97fa09 100644 --- a/core/model/metric/__init__.py +++ b/core/model/metric/__init__.py @@ -13,4 +13,5 @@ from .deepbdc import DeepBDC from .frn import FRN from .meta_baseline import MetaBaseline -from .mcl import MCL \ No newline at end of file +from .mcl import MCL +from .tpn import TPN \ No newline at end of file diff --git a/core/model/metric/tpn.py b/core/model/metric/tpn.py new file mode 100644 index 00000000..a3e98bbf --- /dev/null +++ b/core/model/metric/tpn.py @@ -0,0 +1,174 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F +import numpy as np +from .metric_model import MetricModel + +class RelationNetwork(nn.Module): + """Graph Construction Module""" + def __init__(self): + super(RelationNetwork, self).__init__() + + self.layer1 = nn.Sequential( + nn.Conv2d(64,64,kernel_size=3,padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(kernel_size=2, padding=1)) + self.layer2 = nn.Sequential( + nn.Conv2d(64,1,kernel_size=3,padding=1), + nn.BatchNorm2d(1), + nn.ReLU(), + nn.MaxPool2d(kernel_size=2, padding=1)) + + self.fc3 = nn.Linear(2*2, 8) + self.fc4 = nn.Linear(8, 1) + + self.m0 = nn.MaxPool2d(2) # max-pool without padding + self.m1 = nn.MaxPool2d(2, padding=1) # max-pool with padding + + def forward(self, x, rn): + + x = x.view(-1,64,5,5) + + out = self.layer1(x) + out = self.layer2(out) + # flatten + out = out.view(out.size(0),-1) + out = F.relu(self.fc3(out)) + out = self.fc4(out) # no relu + + out = out.view(out.size(0),-1) # bs*1 + + return out + + +class TPN(MetricModel): + def __init__(self, alpha, **kwargs): + super().__init__(**kwargs) + + self.relation = RelationNetwork() + + if self.rn == 300: + self.alpha = torch.tensor([alpha], requires_grad=False).to(self.device) + elif self.rn == 30: + self.alpha = nn.Parameter(torch.tensor([alpha]).to(self.device), requires_grad=True) + + def labels_to_onehot(self, labels): + batch_size = labels.size(0) + one_hot = torch.zeros(batch_size, self.way_num).to(self.device) + one_hot.scatter_(1, labels.unsqueeze(1), 1) + + return one_hot + + def label_propagation(self, support, query, s_label, q_label): + eps = np.finfo(float).eps + + inp = torch.cat((support, query), 0) + emb_all = self.emb_func(inp).view(-1, 1600) + N, d = emb_all.shape[0], emb_all.shape[1] + + if self.rn in [30, 300]: + self.sigma = self.relation(emb_all, self.rn) + emb_all = emb_all / (self.sigma + eps) + emb1 = torch.unsqueeze(emb_all,1) + emb2 = torch.unsqueeze(emb_all,0) + W = ((emb1-emb2)**2).mean(2) + W = torch.exp(-W/2) + + if self.topk > 0: + topk, indices = torch.topk(W, self.topk) + mask = torch.zeros_like(W) + mask = mask.scatter(1, indices, 1) + mask = ((mask + torch.t(mask)) > 0).type(torch.float32) + W = W * mask + + D = W.sum(0) + D_sqrt_inv = torch.sqrt(1.0 / (D + eps)) + D1 = torch.unsqueeze(D_sqrt_inv, 1).repeat(1, N) + D2 = torch.unsqueeze(D_sqrt_inv, 0).repeat(N, 1) + S = D1 * W * D2 + + ys = s_label + yu = torch.zeros(self.way_num * self.query_num, self.way_num).to(self.device) + y = torch.cat((ys, yu), 0) + F = torch.matmul(torch.inverse(torch.eye(N).to(self.device) - self.alpha * S + eps), y) + Fq = F[self.way_num * self.shot_num:, :] + + gt = torch.argmax(torch.cat((s_label, q_label), 0), 1) + criterion = nn.CrossEntropyLoss() + loss = criterion(F, gt) + + predq = torch.argmax(Fq,1) + gtq = torch.argmax(q_label,1) + correct = (predq==gtq).sum() + total = self.query_num * self.way_num + acc = 1.0 * correct.float() / float(total) + + acc = torch.tensor([acc]).to(self.device) + + return loss, acc + + + def set_forward_loss(self, batch): + image, global_target = batch + image = image.to(self.device) + + episode_size = image.size(0) // (self.way_num * (self.shot_num + self.query_num)) + + ( + support_image, + query_image, + support_target, + query_target, + ) = self.split_by_episode(image, mode=2) + + loss_list = [] + acc_list = [] + + for i in range(episode_size): + s_label_onehot = self.labels_to_onehot(support_target[i]) + q_label_onehot = self.labels_to_onehot(query_target[i]) + + loss, acc = self.label_propagation(support_image[i], query_image[i], s_label_onehot, q_label_onehot) + + loss_list.append(loss) + acc_list.append(acc) + + loss = torch.stack(loss_list) + acc = torch.stack(acc_list) + acc = torch.mean(acc) * 100.0 + + return None, acc, loss + + + def set_forward(self, batch): + image, global_target = batch + image = image.to(self.device) + + episode_size = image.size(0) // (self.way_num * (self.shot_num + self.query_num)) + + ( + support_image, + query_image, + support_target, + query_target, + ) = self.split_by_episode(image, mode=2) + + + acc_list = [] + + for i in range(episode_size): + s_label_onehot = self.labels_to_onehot(support_target[i]) + q_label_onehot = self.labels_to_onehot(query_target[i]) + + _, acc = self.label_propagation(support_image[i], query_image[i], s_label_onehot, q_label_onehot) + + acc_list.append(acc) + + acc = torch.stack(acc_list) + acc = torch.mean(acc) * 100.0 + + return None, acc + + + diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml new file mode 100644 index 00000000..32827a71 --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 1 +query_num: 15 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-10-1-Table1 + +result_root: ./results diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml new file mode 100644 index 00000000..3305fad0 --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 5 +query_num: 15 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-10-5-Table1 + +result_root: ./results diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml new file mode 100644 index 00000000..fbfb5ebd --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 1 +query_num: 15 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-5-1-Table1 + +result_root: ./results diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml new file mode 100644 index 00000000..b90e1140 --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 5 +query_num: 15 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-5-5-Table1 + +result_root: ./results diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml new file mode 100644 index 00000000..8191de01 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 1 +query_num: 15 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-10-1-Table2 + +result_root: ./results diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml new file mode 100644 index 00000000..fe336a14 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 5 +query_num: 15 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-10-5-Table2 + +result_root: ./results diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml new file mode 100644 index 00000000..b59f2644 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 1 +query_num: 15 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-5-1-Table2 + +result_root: ./results diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml new file mode 100644 index 00000000..e57348e7 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml @@ -0,0 +1,26 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 5 +query_num: 15 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-5-5-Table2 + +result_root: ./results From 955e6ea1bb5da084cc850aed8377798af5856f23 Mon Sep 17 00:00:00 2001 From: hanjiang Date: Sun, 13 Jul 2025 13:58:33 +0800 Subject: [PATCH 2/9] =?UTF-8?q?=E2=9C=A8=20feat:=20test=20epoch=20choice?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/model/backbone/__init__.py | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml | 2 +- reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml | 2 +- reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml | 2 +- reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml | 2 +- reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml | 2 +- run_test.py | 1 + 10 files changed, 10 insertions(+), 9 deletions(-) diff --git a/core/model/backbone/__init__.py b/core/model/backbone/__init__.py index 47aa8388..1843d7c3 100644 --- a/core/model/backbone/__init__.py +++ b/core/model/backbone/__init__.py @@ -10,7 +10,7 @@ from .swin_transformer import swin_s, swin_l, swin_b, swin_t, swin_mini from .resnet_bdc import resnet12Bdc, resnet18Bdc from core.model.backbone.utils.maml_module import convert_maml_module - +from .tpn_encoder import CNNEncoder def get_backbone(config): """Get the backbone according to the config dict. diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml index 32827a71..241163db 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-miniImageNet--ravi-10-1-Table1 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml index 3305fad0..f813e9a3 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-miniImageNet--ravi-10-5-Table1 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml index fbfb5ebd..fdcd9dd6 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-miniImageNet--ravi-5-1-Table1 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml index b90e1140..fb8eace4 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-miniImageNet--ravi-5-5-Table1 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml index 8191de01..be97ddc8 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-tieredImageNet-10-1-Table2 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml index fe336a14..d689fc14 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-tieredImageNet-10-5-Table2 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml index b59f2644..f70853e0 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-tieredImageNet-5-1-Table2 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml index e57348e7..ab27e946 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml @@ -18,9 +18,9 @@ seed: 0 n_gpu: 1 device_ids: 0 - log_interval: 100 log_level: info log_name: TPN-tieredImageNet-5-5-Table2 result_root: ./results +save_interval: 100 \ No newline at end of file diff --git a/run_test.py b/run_test.py index 958c87f2..4d5d159b 100644 --- a/run_test.py +++ b/run_test.py @@ -16,6 +16,7 @@ "n_gpu": 2, "test_episode": 600, "episode_size": 2, + "checkpoint_type": "best", # best, last or an epoch number } From 40a86fa73e4dabdd6a2375b43398b94757353651 Mon Sep 17 00:00:00 2001 From: hanjiang Date: Sun, 13 Jul 2025 14:27:40 +0800 Subject: [PATCH 3/9] =?UTF-8?q?=F0=9F=90=9E=20fix:=20dataset=20config?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config/TPN.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml | 2 +- reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml | 2 +- reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml | 8 +++++++- reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml | 8 +++++++- reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml | 8 +++++++- reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml | 8 +++++++- 9 files changed, 33 insertions(+), 9 deletions(-) diff --git a/config/TPN.yaml b/config/TPN.yaml index 63dfcba5..7da58273 100644 --- a/config/TPN.yaml +++ b/config/TPN.yaml @@ -29,4 +29,4 @@ lr_scheduler: name: StepLR kwargs: step_size: 10000 - gamma: 0.5 \ No newline at end of file + gamma: 0.5 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml index 241163db..b3dd7d12 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-Table1.yaml @@ -23,4 +23,4 @@ log_level: info log_name: TPN-miniImageNet--ravi-10-1-Table1 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml index f813e9a3..b7ff1e8d 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-Table1.yaml @@ -23,4 +23,4 @@ log_level: info log_name: TPN-miniImageNet--ravi-10-5-Table1 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml index fdcd9dd6..723fd6d5 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-Table1.yaml @@ -23,4 +23,4 @@ log_level: info log_name: TPN-miniImageNet--ravi-5-1-Table1 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml index fb8eace4..c595f518 100644 --- a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-Table1.yaml @@ -23,4 +23,4 @@ log_level: info log_name: TPN-miniImageNet--ravi-5-5-Table1 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml index be97ddc8..2142c735 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml @@ -6,6 +6,12 @@ includes: - headers/optimizer.yaml - TPN.yaml +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + way_num: 10 shot_num: 1 query_num: 15 @@ -23,4 +29,4 @@ log_level: info log_name: TPN-tieredImageNet-10-1-Table2 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml index d689fc14..1d4df020 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml @@ -6,6 +6,12 @@ includes: - headers/optimizer.yaml - TPN.yaml +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + way_num: 10 shot_num: 5 query_num: 15 @@ -23,4 +29,4 @@ log_level: info log_name: TPN-tieredImageNet-10-5-Table2 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml index f70853e0..dc8640b2 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml @@ -6,6 +6,12 @@ includes: - headers/optimizer.yaml - TPN.yaml +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + way_num: 5 shot_num: 1 query_num: 15 @@ -23,4 +29,4 @@ log_level: info log_name: TPN-tieredImageNet-5-1-Table2 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml index ab27e946..61700646 100644 --- a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml +++ b/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml @@ -6,6 +6,12 @@ includes: - headers/optimizer.yaml - TPN.yaml +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + way_num: 5 shot_num: 5 query_num: 15 @@ -23,4 +29,4 @@ log_level: info log_name: TPN-tieredImageNet-5-5-Table2 result_root: ./results -save_interval: 100 \ No newline at end of file +save_interval: 100 From f1f21ca38f37a47ed5d72d2141e6ae1b2959e445 Mon Sep 17 00:00:00 2001 From: hanjiang Date: Sun, 13 Jul 2025 14:36:14 +0800 Subject: [PATCH 4/9] =?UTF-8?q?=F0=9F=90=9E=20fix:=20config=20name?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ...-ravi-10-1-Table2.yaml => TPN-tieredImageNet-10-1-Table2.yaml} | 0 ...-ravi-10-5-Table2.yaml => TPN-tieredImageNet-10-5-Table2.yaml} | 0 ...t--ravi-5-1-Table2.yaml => TPN-tieredImageNet-5-1-Table2.yaml} | 0 ...t--ravi-5-5-Table2.yaml => TPN-tieredImageNet-5-5-Table2.yaml} | 0 4 files changed, 0 insertions(+), 0 deletions(-) rename reproduce/TPN/{TPN-tieredImageNet--ravi-10-1-Table2.yaml => TPN-tieredImageNet-10-1-Table2.yaml} (100%) rename reproduce/TPN/{TPN-tieredImageNet--ravi-10-5-Table2.yaml => TPN-tieredImageNet-10-5-Table2.yaml} (100%) rename reproduce/TPN/{TPN-tieredImageNet--ravi-5-1-Table2.yaml => TPN-tieredImageNet-5-1-Table2.yaml} (100%) rename reproduce/TPN/{TPN-tieredImageNet--ravi-5-5-Table2.yaml => TPN-tieredImageNet-5-5-Table2.yaml} (100%) diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-10-1-Table2.yaml similarity index 100% rename from reproduce/TPN/TPN-tieredImageNet--ravi-10-1-Table2.yaml rename to reproduce/TPN/TPN-tieredImageNet-10-1-Table2.yaml diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-10-5-Table2.yaml similarity index 100% rename from reproduce/TPN/TPN-tieredImageNet--ravi-10-5-Table2.yaml rename to reproduce/TPN/TPN-tieredImageNet-10-5-Table2.yaml diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-5-1-Table2.yaml similarity index 100% rename from reproduce/TPN/TPN-tieredImageNet--ravi-5-1-Table2.yaml rename to reproduce/TPN/TPN-tieredImageNet-5-1-Table2.yaml diff --git a/reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-5-5-Table2.yaml similarity index 100% rename from reproduce/TPN/TPN-tieredImageNet--ravi-5-5-Table2.yaml rename to reproduce/TPN/TPN-tieredImageNet-5-5-Table2.yaml From 7b307982d6abb3b4823875ee9b9d902d6f951281 Mon Sep 17 00:00:00 2001 From: hanjiang Date: Mon, 14 Jul 2025 19:05:34 +0800 Subject: [PATCH 5/9] =?UTF-8?q?=F0=9F=90=9E=20fix:=20loss=20computation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config/TPN.yaml | 4 ++-- core/model/metric/tpn.py | 3 ++- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/config/TPN.yaml b/config/TPN.yaml index 7da58273..77ba66fd 100644 --- a/config/TPN.yaml +++ b/config/TPN.yaml @@ -14,11 +14,11 @@ way_num: 5 shot_num: 1 query_num: 15 -epoch: 2100 +epoch: 1000 test_epoch: 100 train_episode: 100 test_episode: 100 -episode_size: 1 +episode_size: 5 optimizer: name: Adam diff --git a/core/model/metric/tpn.py b/core/model/metric/tpn.py index a3e98bbf..02b82850 100644 --- a/core/model/metric/tpn.py +++ b/core/model/metric/tpn.py @@ -104,7 +104,7 @@ def label_propagation(self, support, query, s_label, q_label): total = self.query_num * self.way_num acc = 1.0 * correct.float() / float(total) - acc = torch.tensor([acc]).to(self.device) + acc = torch.tensor([acc]) return loss, acc @@ -135,6 +135,7 @@ def set_forward_loss(self, batch): acc_list.append(acc) loss = torch.stack(loss_list) + loss = torch.mean(loss) acc = torch.stack(acc_list) acc = torch.mean(acc) * 100.0 From 4ca1095533a91a0134ae33190e6d8eec7b4e284b Mon Sep 17 00:00:00 2001 From: hanjiang Date: Mon, 14 Jul 2025 21:06:29 +0800 Subject: [PATCH 6/9] =?UTF-8?q?=F0=9F=90=9E=20fix:=20test=20epoch=20choice?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/test.py | 17 ++++++++++++++++- 1 file changed, 16 insertions(+), 1 deletion(-) diff --git a/core/test.py b/core/test.py index 0da3167a..79b78e23 100644 --- a/core/test.py +++ b/core/test.py @@ -193,7 +193,22 @@ def _init_files(self, config): rank=self.rank, ) - state_dict_path = os.path.join(result_path, "checkpoints", "model_best.pth") + checkpoint_type = config.get("checkpoint_type", "best") + if checkpoint_type == "best": + checkpoint_filename = "model_best.pth" + elif checkpoint_type == "last": + checkpoint_filename = "model_last.pth" + elif isinstance(checkpoint_type, int) or checkpoint_type.isdigit(): + epoch_num = int(checkpoint_type) + checkpoint_filename = f"model_{epoch_num:05d}.pth" + else: + print( + f"Warning: Invalid checkpoint_type '{checkpoint_type}', using 'best'", + level="warning", + ) + checkpoint_filename = "model_best.pth" + + state_dict_path = os.path.join(result_path, "checkpoints", checkpoint_filename) if self.rank == 0: create_dirs([result_path, log_path, viz_path]) From c59c862d82c066efd0d2874b22d77d3d87d7188a Mon Sep 17 00:00:00 2001 From: hanjiang Date: Tue, 15 Jul 2025 10:19:23 +0800 Subject: [PATCH 7/9] =?UTF-8?q?=F0=9F=A6=84=20refactor:=20update=20README?= =?UTF-8?q?=20&=20rm=20unnecessary=20implementation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- config/TPN.yaml | 11 +++++-- core/model/backbone/__init__.py | 1 - core/model/backbone/tpn_encoder.py | 49 ------------------------------ core/model/metric/tpn.py | 14 +++++++++ reproduce/TPN/README.md | 36 ++++++++++++++++++++++ 5 files changed, 59 insertions(+), 52 deletions(-) delete mode 100644 core/model/backbone/tpn_encoder.py create mode 100644 reproduce/TPN/README.md diff --git a/config/TPN.yaml b/config/TPN.yaml index 77ba66fd..9c4ff735 100644 --- a/config/TPN.yaml +++ b/config/TPN.yaml @@ -1,6 +1,13 @@ backbone: - name: CNNEncoder - kwargs: null + name: Conv64F + kwargs: + is_flatten: false + is_feature: false + leaky_relu: false + negative_slope: 0.2 + last_pool: true + maxpool_last2: true + use_running_statistics: true classifier: name: TPN diff --git a/core/model/backbone/__init__.py b/core/model/backbone/__init__.py index 1843d7c3..19077ab9 100644 --- a/core/model/backbone/__init__.py +++ b/core/model/backbone/__init__.py @@ -10,7 +10,6 @@ from .swin_transformer import swin_s, swin_l, swin_b, swin_t, swin_mini from .resnet_bdc import resnet12Bdc, resnet18Bdc from core.model.backbone.utils.maml_module import convert_maml_module -from .tpn_encoder import CNNEncoder def get_backbone(config): """Get the backbone according to the config dict. diff --git a/core/model/backbone/tpn_encoder.py b/core/model/backbone/tpn_encoder.py deleted file mode 100644 index 42f9c2b0..00000000 --- a/core/model/backbone/tpn_encoder.py +++ /dev/null @@ -1,49 +0,0 @@ -#------------------------------------- -# Project: Transductive Propagation Network for Few-shot Learning -# Date: 2019.1.11 -# Author: Yanbin Liu -# All Rights Reserved -#------------------------------------- - -import torch -import torch.nn as nn - -class CNNEncoder(nn.Module): - """Encoder for feature embedding""" - def __init__(self): - super(CNNEncoder, self).__init__() - self.layer1 = nn.Sequential( - nn.Conv2d(3, 64, kernel_size=3, padding=1), - nn.BatchNorm2d(64), - nn.ReLU(), - nn.MaxPool2d(2)) - self.layer2 = nn.Sequential( - nn.Conv2d(64,64,kernel_size=3,padding=1), - nn.BatchNorm2d(64), - nn.ReLU(), - nn.MaxPool2d(2)) - self.layer3 = nn.Sequential( - nn.Conv2d(64,64,kernel_size=3,padding=1), - nn.BatchNorm2d(64), - nn.ReLU(), - nn.MaxPool2d(2)) - self.layer4 = nn.Sequential( - nn.Conv2d(64,64,kernel_size=3,padding=1), - nn.BatchNorm2d(64), - nn.ReLU(), - nn.MaxPool2d(2)) - - def forward(self,x): - """x: bs*3*84*84 """ - out = self.layer1(x) - out = self.layer2(out) - out = self.layer3(out) - out = self.layer4(out) - - return out - - - - - - diff --git a/core/model/metric/tpn.py b/core/model/metric/tpn.py index 02b82850..7dfb9852 100644 --- a/core/model/metric/tpn.py +++ b/core/model/metric/tpn.py @@ -1,3 +1,17 @@ +""" +@misc{liu2019learningpropagatelabelstransductive, + title={Learning to Propagate Labels: Transductive Propagation Network for Few-shot Learning}, + author={Yanbin Liu and Juho Lee and Minseop Park and Saehoon Kim and Eunho Yang and Sung Ju Hwang and Yi Yang}, + year={2019}, + eprint={1805.10002}, + archivePrefix={arXiv}, + primaryClass={cs.LG}, + url={https://arxiv.org/abs/1805.10002}, +} + +Adapted From https://github.com/csyanbin/TPN-pytorch +""" + import torch import torch.nn as nn import torch.nn.functional as F diff --git a/reproduce/TPN/README.md b/reproduce/TPN/README.md new file mode 100644 index 00000000..dd619a91 --- /dev/null +++ b/reproduce/TPN/README.md @@ -0,0 +1,36 @@ +# TPN Reproduction + +## Introduction + +| Name: | [TPN](https://arxiv.org/abs/1805.10002) | +| ------- | ---------------------------------------------------------- | +| Embed.: | Conv64F | +| Type: | Metric | +| Venue: | ICLR2019 | +| Codes: | [**TPN-pytorch**](https://github.com/csyanbin/TPN-pytorch) | + +Cite this work with (template): + +```bibtex +@misc{liu2019learningpropagatelabelstransductive, + title={Learning to Propagate Labels: Transductive Propagation Network for Few-shot Learning}, + author={Yanbin Liu and Juho Lee and Minseop Park and Saehoon Kim and Eunho Yang and Sung Ju Hwang and Yi Yang}, + year={2019}, + eprint={1805.10002}, + archivePrefix={arXiv}, + primaryClass={cs.LG}, + url={https://arxiv.org/abs/1805.10002}, +} +``` + +--- + +## Results and Models + +All the results are tested under the best model. Checkpoints of different epochs are also provided. + +| dataset/task | 5way-1shot | 5way-5shot | 10way-1shot | 10-way-5shot | +| -------------- | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | +| miniImageNet | 54.06±0.38 [:arrow_down:](https://drive.google.com/drive/folders/1Y14e-h_DcwfyxwU71GZIXQ2G39ZS8oys) | 69.27±0.30 [:arrow_down:](https://drive.google.com/drive/folders/1DIBmJ8a_GZIlEmUW0KEf_awTaB-8GB7i) | 37.36±0.23 [:arrow_down:](https://drive.google.com/drive/folders/1OeO3K7wY4y-UN979eRUQo1vvVvin3vF-) | 53.62±0.20 [:arrow_down:](https://drive.google.com/drive/folders/1c8yd0rMQhAytePcrfLq2nEYaxPDFHc8f) | +| tieredImageNet | 53.36±0.42 [:arrow_down:](https://drive.google.com/drive/folders/1_C0VA1LirJ5l3kYqEBY8HfHmXOzCkZog) | 69.83±0.35 [:arrow_down:](https://drive.google.com/drive/folders/1anwG8tjvaXQ5oq9BBjcTQjY9OdKmDYf1) | 40.29±0.28 [:arrow_down:](https://drive.google.com/drive/folders/1HnJfDHHE78YzaGvhmHHvaG4ymVhdPkrh) | 57.53±0.25 [:arrow_down:](https://drive.google.com/drive/folders/1NBaxyY60rJIxnoeOV0UAAt6KORPvZyfJ) | + From e84f57e2053ff95cb82691a6276225a9ca51b4fc Mon Sep 17 00:00:00 2001 From: hanjiang Date: Wed, 23 Jul 2025 13:45:07 +0800 Subject: [PATCH 8/9] =?UTF-8?q?=F0=9F=A6=84=20refactor:=20update=20more=20?= =?UTF-8?q?results=20&=20rename=20variables?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/model/metric/tpn.py | 175 ++++++++++-------- reproduce/TPN/README.md | 9 + ...ImageNet--ravi-10-1-highershot-Table1.yaml | 27 +++ ...ImageNet--ravi-10-5-highershot-Table1.yaml | 27 +++ ...iImageNet--ravi-5-1-highershot-Table1.yaml | 27 +++ ...iImageNet--ravi-5-5-highershot-Table1.yaml | 27 +++ ...tieredImageNet-10-1-highershot-Table2.yaml | 33 ++++ ...tieredImageNet-10-5-highershot-Table2.yaml | 33 ++++ ...-tieredImageNet-5-1-highershot-Table2.yaml | 33 ++++ ...-tieredImageNet-5-5-highershot-Table2.yaml | 33 ++++ 10 files changed, 348 insertions(+), 76 deletions(-) create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-10-1-highershot-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-10-5-highershot-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-5-1-highershot-Table1.yaml create mode 100644 reproduce/TPN/TPN-miniImageNet--ravi-5-5-highershot-Table1.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet-10-1-highershot-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet-10-5-highershot-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet-5-1-highershot-Table2.yaml create mode 100644 reproduce/TPN/TPN-tieredImageNet-5-5-highershot-Table2.yaml diff --git a/core/model/metric/tpn.py b/core/model/metric/tpn.py index 7dfb9852..01411c21 100644 --- a/core/model/metric/tpn.py +++ b/core/model/metric/tpn.py @@ -1,57 +1,61 @@ """ @misc{liu2019learningpropagatelabelstransductive, - title={Learning to Propagate Labels: Transductive Propagation Network for Few-shot Learning}, + title={Learning to Propagate Labels: Transductive Propagation Network for Few-shot Learning}, author={Yanbin Liu and Juho Lee and Minseop Park and Saehoon Kim and Eunho Yang and Sung Ju Hwang and Yi Yang}, year={2019}, eprint={1805.10002}, archivePrefix={arXiv}, primaryClass={cs.LG}, - url={https://arxiv.org/abs/1805.10002}, + url={https://arxiv.org/abs/1805.10002}, } Adapted From https://github.com/csyanbin/TPN-pytorch """ +import numpy as np import torch import torch.nn as nn import torch.nn.functional as F -import numpy as np + from .metric_model import MetricModel + class RelationNetwork(nn.Module): """Graph Construction Module""" + def __init__(self): super(RelationNetwork, self).__init__() self.layer1 = nn.Sequential( - nn.Conv2d(64,64,kernel_size=3,padding=1), - nn.BatchNorm2d(64), - nn.ReLU(), - nn.MaxPool2d(kernel_size=2, padding=1)) + nn.Conv2d(64, 64, kernel_size=3, padding=1), + nn.BatchNorm2d(64), + nn.ReLU(), + nn.MaxPool2d(kernel_size=2, padding=1), + ) self.layer2 = nn.Sequential( - nn.Conv2d(64,1,kernel_size=3,padding=1), - nn.BatchNorm2d(1), - nn.ReLU(), - nn.MaxPool2d(kernel_size=2, padding=1)) + nn.Conv2d(64, 1, kernel_size=3, padding=1), + nn.BatchNorm2d(1), + nn.ReLU(), + nn.MaxPool2d(kernel_size=2, padding=1), + ) - self.fc3 = nn.Linear(2*2, 8) + self.fc3 = nn.Linear(2 * 2, 8) self.fc4 = nn.Linear(8, 1) - self.m0 = nn.MaxPool2d(2) # max-pool without padding - self.m1 = nn.MaxPool2d(2, padding=1) # max-pool with padding + self.m0 = nn.MaxPool2d(2) # max-pool without padding + self.m1 = nn.MaxPool2d(2, padding=1) # max-pool with padding def forward(self, x, rn): - - x = x.view(-1,64,5,5) - + x = x.view(-1, 64, 5, 5) + out = self.layer1(x) out = self.layer2(out) # flatten - out = out.view(out.size(0),-1) + out = out.view(out.size(0), -1) out = F.relu(self.fc3(out)) - out = self.fc4(out) # no relu + out = self.fc4(out) # no relu - out = out.view(out.size(0),-1) # bs*1 + out = out.view(out.size(0), -1) # bs*1 return out @@ -65,7 +69,9 @@ def __init__(self, alpha, **kwargs): if self.rn == 300: self.alpha = torch.tensor([alpha], requires_grad=False).to(self.device) elif self.rn == 30: - self.alpha = nn.Parameter(torch.tensor([alpha]).to(self.device), requires_grad=True) + self.alpha = nn.Parameter( + torch.tensor([alpha]).to(self.device), requires_grad=True + ) def labels_to_onehot(self, labels): batch_size = labels.size(0) @@ -74,60 +80,70 @@ def labels_to_onehot(self, labels): return one_hot - def label_propagation(self, support, query, s_label, q_label): + def label_propagation(self, support, query, support_label, query_label): eps = np.finfo(float).eps - inp = torch.cat((support, query), 0) - emb_all = self.emb_func(inp).view(-1, 1600) - N, d = emb_all.shape[0], emb_all.shape[1] + input_feat = torch.cat((support, query), 0) + embedding_all = self.emb_func(input_feat).view(-1, 1600) + num_nodes = embedding_all.shape[0] if self.rn in [30, 300]: - self.sigma = self.relation(emb_all, self.rn) - emb_all = emb_all / (self.sigma + eps) - emb1 = torch.unsqueeze(emb_all,1) - emb2 = torch.unsqueeze(emb_all,0) - W = ((emb1-emb2)**2).mean(2) - W = torch.exp(-W/2) + self.sigma = self.relation(embedding_all, self.rn) + embedding_all = embedding_all / (self.sigma + eps) + embedding_1 = torch.unsqueeze(embedding_all, 1) + embedding_2 = torch.unsqueeze(embedding_all, 0) + weight_matrix = ((embedding_1 - embedding_2) ** 2).mean(2) + weight_matrix = torch.exp(-weight_matrix / 2) if self.topk > 0: - topk, indices = torch.topk(W, self.topk) - mask = torch.zeros_like(W) - mask = mask.scatter(1, indices, 1) + topk_values, topk_indices = torch.topk(weight_matrix, self.topk) + mask = torch.zeros_like(weight_matrix) + mask = mask.scatter(1, topk_indices, 1) mask = ((mask + torch.t(mask)) > 0).type(torch.float32) - W = W * mask - - D = W.sum(0) - D_sqrt_inv = torch.sqrt(1.0 / (D + eps)) - D1 = torch.unsqueeze(D_sqrt_inv, 1).repeat(1, N) - D2 = torch.unsqueeze(D_sqrt_inv, 0).repeat(N, 1) - S = D1 * W * D2 - - ys = s_label - yu = torch.zeros(self.way_num * self.query_num, self.way_num).to(self.device) - y = torch.cat((ys, yu), 0) - F = torch.matmul(torch.inverse(torch.eye(N).to(self.device) - self.alpha * S + eps), y) - Fq = F[self.way_num * self.shot_num:, :] - - gt = torch.argmax(torch.cat((s_label, q_label), 0), 1) + weight_matrix = weight_matrix * mask + + degree_matrix = weight_matrix.sum(0) + degree_sqrt_inv = torch.sqrt(1.0 / (degree_matrix + eps)) + degree_1 = torch.unsqueeze(degree_sqrt_inv, 1).repeat(1, num_nodes) + degree_2 = torch.unsqueeze(degree_sqrt_inv, 0).repeat(num_nodes, 1) + symmetric_matrix = degree_1 * weight_matrix * degree_2 + + support_labels = support_label + unlabeled_query = torch.zeros(self.way_num * self.query_num, self.way_num).to( + self.device + ) + combined_labels = torch.cat((support_labels, unlabeled_query), 0) + propagated_labels = torch.matmul( + torch.inverse( + torch.eye(num_nodes).to(self.device) + - self.alpha * symmetric_matrix + + eps + ), + combined_labels, + ) + query_predictions = propagated_labels[self.way_num * self.shot_num :, :] + + ground_truth = torch.argmax(torch.cat((support_label, query_label), 0), 1) criterion = nn.CrossEntropyLoss() - loss = criterion(F, gt) + loss = criterion(propagated_labels, ground_truth) - predq = torch.argmax(Fq,1) - gtq = torch.argmax(q_label,1) - correct = (predq==gtq).sum() - total = self.query_num * self.way_num - acc = 1.0 * correct.float() / float(total) + predicted_query = torch.argmax(query_predictions, 1) + ground_truth_query = torch.argmax(query_label, 1) + correct_predictions = (predicted_query == ground_truth_query).sum() + total_queries = self.query_num * self.way_num + accuracy = 1.0 * correct_predictions.float() / float(total_queries) - acc = torch.tensor([acc]) - - return loss, acc + accuracy = torch.tensor([accuracy]) + return loss, accuracy def set_forward_loss(self, batch): image, global_target = batch image = image.to(self.device) - episode_size = image.size(0) // (self.way_num * (self.shot_num + self.query_num)) + episode_size = image.size(0) // ( + self.way_num * (self.shot_num + self.query_num) + ) ( support_image, @@ -140,27 +156,33 @@ def set_forward_loss(self, batch): acc_list = [] for i in range(episode_size): - s_label_onehot = self.labels_to_onehot(support_target[i]) - q_label_onehot = self.labels_to_onehot(query_target[i]) - - loss, acc = self.label_propagation(support_image[i], query_image[i], s_label_onehot, q_label_onehot) + support_label_onehot = self.labels_to_onehot(support_target[i]) + query_label_onehot = self.labels_to_onehot(query_target[i]) + + loss, acc = self.label_propagation( + support_image[i], + query_image[i], + support_label_onehot, + query_label_onehot, + ) loss_list.append(loss) acc_list.append(acc) - + loss = torch.stack(loss_list) - loss = torch.mean(loss) + loss = torch.mean(loss) acc = torch.stack(acc_list) acc = torch.mean(acc) * 100.0 return None, acc, loss - def set_forward(self, batch): image, global_target = batch image = image.to(self.device) - episode_size = image.size(0) // (self.way_num * (self.shot_num + self.query_num)) + episode_size = image.size(0) // ( + self.way_num * (self.shot_num + self.query_num) + ) ( support_image, @@ -169,21 +191,22 @@ def set_forward(self, batch): query_target, ) = self.split_by_episode(image, mode=2) - acc_list = [] for i in range(episode_size): - s_label_onehot = self.labels_to_onehot(support_target[i]) - q_label_onehot = self.labels_to_onehot(query_target[i]) - - _, acc = self.label_propagation(support_image[i], query_image[i], s_label_onehot, q_label_onehot) + support_label_onehot = self.labels_to_onehot(support_target[i]) + query_label_onehot = self.labels_to_onehot(query_target[i]) + + _, acc = self.label_propagation( + support_image[i], + query_image[i], + support_label_onehot, + query_label_onehot, + ) acc_list.append(acc) - + acc = torch.stack(acc_list) acc = torch.mean(acc) * 100.0 return None, acc - - - diff --git a/reproduce/TPN/README.md b/reproduce/TPN/README.md index dd619a91..367c5349 100644 --- a/reproduce/TPN/README.md +++ b/reproduce/TPN/README.md @@ -29,8 +29,17 @@ Cite this work with (template): All the results are tested under the best model. Checkpoints of different epochs are also provided. +**TPN Result** + | dataset/task | 5way-1shot | 5way-5shot | 10way-1shot | 10-way-5shot | | -------------- | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | | miniImageNet | 54.06±0.38 [:arrow_down:](https://drive.google.com/drive/folders/1Y14e-h_DcwfyxwU71GZIXQ2G39ZS8oys) | 69.27±0.30 [:arrow_down:](https://drive.google.com/drive/folders/1DIBmJ8a_GZIlEmUW0KEf_awTaB-8GB7i) | 37.36±0.23 [:arrow_down:](https://drive.google.com/drive/folders/1OeO3K7wY4y-UN979eRUQo1vvVvin3vF-) | 53.62±0.20 [:arrow_down:](https://drive.google.com/drive/folders/1c8yd0rMQhAytePcrfLq2nEYaxPDFHc8f) | | tieredImageNet | 53.36±0.42 [:arrow_down:](https://drive.google.com/drive/folders/1_C0VA1LirJ5l3kYqEBY8HfHmXOzCkZog) | 69.83±0.35 [:arrow_down:](https://drive.google.com/drive/folders/1anwG8tjvaXQ5oq9BBjcTQjY9OdKmDYf1) | 40.29±0.28 [:arrow_down:](https://drive.google.com/drive/folders/1HnJfDHHE78YzaGvhmHHvaG4ymVhdPkrh) | 57.53±0.25 [:arrow_down:](https://drive.google.com/drive/folders/1NBaxyY60rJIxnoeOV0UAAt6KORPvZyfJ) | +**Higher Shot TPN Result** + +| dataset/task | 5way-1shot | 5way-5shot | 10way-1shot | 10-way-5shot | +| -------------- | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | ------------------------------------------------------------ | +| miniImageNet | 55.37±0.39 [:arrow_down:](https://drive.google.com/drive/folders/1sMAcvy817oBzMybkO_HVQ0AMBhJW0kbr?usp=drive_link) | 68.80±0.30 [:arrow_down:](https://drive.google.com/drive/folders/1G1WMnHbsgJcgbdSWoVqRWkQJldDvbDL1?usp=drive_link) | 38.62±0.22 [:arrow_down:](https://drive.google.com/drive/folders/1aXKD2GBbsM7Ql0FdpZV1NJeEdu5QZMyV?usp=drive_link) | 53.68±0.21 [:arrow_down:](https://drive.google.com/drive/folders/1mi2r9Kbl6SvSdbn5kKRd3DaqS354vsxe?usp=drive_link) | +| tieredImageNet | 57.17±0.42 [:arrow_down:](https://drive.google.com/drive/folders/1G-zBq1Hp0zrPIgN1UUlOIeuXnusRJfmB?usp=drive_link) | 69.67±0.34 [:arrow_down:](https://drive.google.com/drive/folders/1-1e22Tjs1YtVuh6HujQ5hMCXWORlKVBR?usp=drive_link) | 43.40±0.28 [:arrow_down:](https://drive.google.com/drive/folders/1hY0q3Nd84SqBSTVmF1ZveKkxc5ZH5raf?usp=drive_link) | 57.71±0.25 [:arrow_down:](https://drive.google.com/drive/folders/1bwoHxwTBSBtPHOyv-0B1v2ZS2IMbAlO1?usp=drive_link) | + diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-1-highershot-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-highershot-Table1.yaml new file mode 100644 index 00000000..44f6c676 --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-1-highershot-Table1.yaml @@ -0,0 +1,27 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 5 +query_num: 15 +test_shot: 1 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-10-1-highershot-Table1 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-10-5-highershot-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-highershot-Table1.yaml new file mode 100644 index 00000000..9d40db2e --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-10-5-highershot-Table1.yaml @@ -0,0 +1,27 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 10 +shot_num: 10 +query_num: 15 +test_shot: 5 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-10-5-highershot-Table1 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-1-highershot-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-highershot-Table1.yaml new file mode 100644 index 00000000..f68c56f0 --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-1-highershot-Table1.yaml @@ -0,0 +1,27 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 5 +query_num: 15 +test_shot: 1 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-5-1-highershot-Table1 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-miniImageNet--ravi-5-5-highershot-Table1.yaml b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-highershot-Table1.yaml new file mode 100644 index 00000000..569ee0eb --- /dev/null +++ b/reproduce/TPN/TPN-miniImageNet--ravi-5-5-highershot-Table1.yaml @@ -0,0 +1,27 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +way_num: 5 +shot_num: 10 +query_num: 15 +test_shot: 5 + +data_root: /data/fewshot/miniImageNet--ravi +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-miniImageNet--ravi-5-5-highershot-Table1 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet-10-1-highershot-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-10-1-highershot-Table2.yaml new file mode 100644 index 00000000..94764a1f --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet-10-1-highershot-Table2.yaml @@ -0,0 +1,33 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + +way_num: 10 +shot_num: 5 +query_num: 15 +test_shot: 1 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-10-1-highershot-Table2 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet-10-5-highershot-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-10-5-highershot-Table2.yaml new file mode 100644 index 00000000..500866c7 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet-10-5-highershot-Table2.yaml @@ -0,0 +1,33 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + +way_num: 10 +shot_num: 10 +query_num: 15 +test_shot: 5 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-10-5-highershot-Table2 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet-5-1-highershot-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-5-1-highershot-Table2.yaml new file mode 100644 index 00000000..f5851483 --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet-5-1-highershot-Table2.yaml @@ -0,0 +1,33 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + +way_num: 5 +shot_num: 5 +query_num: 15 +test_shot: 1 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-5-1-highershot-Table2 + +result_root: ./results +save_interval: 100 diff --git a/reproduce/TPN/TPN-tieredImageNet-5-5-highershot-Table2.yaml b/reproduce/TPN/TPN-tieredImageNet-5-5-highershot-Table2.yaml new file mode 100644 index 00000000..ffce4d1e --- /dev/null +++ b/reproduce/TPN/TPN-tieredImageNet-5-5-highershot-Table2.yaml @@ -0,0 +1,33 @@ +includes: + - headers/data.yaml + - headers/device.yaml + - headers/misc.yaml + - headers/model.yaml + - headers/optimizer.yaml + - TPN.yaml + +lr_scheduler: + name: StepLR + kwargs: + step_size: 25000 + gamma: 0.5 + +way_num: 5 +shot_num: 10 +query_num: 15 +test_shot: 5 + +data_root: /data/fewshot/tiered_imagenet +use_memory: false + +seed: 0 + +n_gpu: 1 +device_ids: 0 + +log_interval: 100 +log_level: info +log_name: TPN-tieredImageNet-5-5-highershot-Table2 + +result_root: ./results +save_interval: 100 From 3cee06002aa235452c97a8785f1ffb06aabb0d40 Mon Sep 17 00:00:00 2001 From: hanjiang Date: Thu, 24 Jul 2025 22:15:54 +0800 Subject: [PATCH 9/9] =?UTF-8?q?=F0=9F=8E=88=20perf:=20improve=20code=20per?= =?UTF-8?q?f?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- core/model/metric/tpn.py | 85 +++++++++++++++++----------------------- 1 file changed, 35 insertions(+), 50 deletions(-) diff --git a/core/model/metric/tpn.py b/core/model/metric/tpn.py index 01411c21..c7c62aa5 100644 --- a/core/model/metric/tpn.py +++ b/core/model/metric/tpn.py @@ -42,20 +42,17 @@ def __init__(self): self.fc3 = nn.Linear(2 * 2, 8) self.fc4 = nn.Linear(8, 1) - self.m0 = nn.MaxPool2d(2) # max-pool without padding - self.m1 = nn.MaxPool2d(2, padding=1) # max-pool with padding - - def forward(self, x, rn): + def forward(self, x): x = x.view(-1, 64, 5, 5) out = self.layer1(x) out = self.layer2(out) - # flatten + out = out.view(out.size(0), -1) out = F.relu(self.fc3(out)) - out = self.fc4(out) # no relu + out = self.fc4(out) - out = out.view(out.size(0), -1) # bs*1 + out = out.view(out.size(0), -1) return out @@ -65,6 +62,7 @@ def __init__(self, alpha, **kwargs): super().__init__(**kwargs) self.relation = RelationNetwork() + self.eps = torch.finfo(torch.float32).eps if self.rn == 300: self.alpha = torch.tensor([alpha], requires_grad=False).to(self.device) @@ -81,59 +79,49 @@ def labels_to_onehot(self, labels): return one_hot def label_propagation(self, support, query, support_label, query_label): - eps = np.finfo(float).eps - input_feat = torch.cat((support, query), 0) embedding_all = self.emb_func(input_feat).view(-1, 1600) num_nodes = embedding_all.shape[0] if self.rn in [30, 300]: - self.sigma = self.relation(embedding_all, self.rn) - embedding_all = embedding_all / (self.sigma + eps) - embedding_1 = torch.unsqueeze(embedding_all, 1) - embedding_2 = torch.unsqueeze(embedding_all, 0) - weight_matrix = ((embedding_1 - embedding_2) ** 2).mean(2) + self.sigma = self.relation(embedding_all) + embedding_all = embedding_all / (self.sigma + self.eps) + + weight_matrix = torch.cdist(embedding_all, embedding_all, p=2) ** 2 weight_matrix = torch.exp(-weight_matrix / 2) if self.topk > 0: topk_values, topk_indices = torch.topk(weight_matrix, self.topk) mask = torch.zeros_like(weight_matrix) mask = mask.scatter(1, topk_indices, 1) - mask = ((mask + torch.t(mask)) > 0).type(torch.float32) + mask = (mask + mask.t()) > 0 weight_matrix = weight_matrix * mask degree_matrix = weight_matrix.sum(0) - degree_sqrt_inv = torch.sqrt(1.0 / (degree_matrix + eps)) - degree_1 = torch.unsqueeze(degree_sqrt_inv, 1).repeat(1, num_nodes) - degree_2 = torch.unsqueeze(degree_sqrt_inv, 0).repeat(num_nodes, 1) - symmetric_matrix = degree_1 * weight_matrix * degree_2 + degree_sqrt_inv = torch.rsqrt(degree_matrix + self.eps) + + symmetric_matrix = weight_matrix * degree_sqrt_inv.unsqueeze(0) * degree_sqrt_inv.unsqueeze(1) support_labels = support_label - unlabeled_query = torch.zeros(self.way_num * self.query_num, self.way_num).to( - self.device - ) + unlabeled_query = torch.zeros(self.way_num * self.query_num, self.way_num, + device=self.device, dtype=support_labels.dtype) combined_labels = torch.cat((support_labels, unlabeled_query), 0) - propagated_labels = torch.matmul( - torch.inverse( - torch.eye(num_nodes).to(self.device) - - self.alpha * symmetric_matrix - + eps - ), - combined_labels, - ) - query_predictions = propagated_labels[self.way_num * self.shot_num :, :] - - ground_truth = torch.argmax(torch.cat((support_label, query_label), 0), 1) + + identity_matrix = torch.eye(num_nodes, device=self.device) + A = identity_matrix - self.alpha * symmetric_matrix + propagated_labels = torch.linalg.solve(A, combined_labels) + + query_predictions = propagated_labels[self.way_num * self.shot_num:, :] + + all_labels = torch.cat((support_label, query_label), 0) + ground_truth = torch.argmax(all_labels, 1) criterion = nn.CrossEntropyLoss() loss = criterion(propagated_labels, ground_truth) predicted_query = torch.argmax(query_predictions, 1) ground_truth_query = torch.argmax(query_label, 1) - correct_predictions = (predicted_query == ground_truth_query).sum() - total_queries = self.query_num * self.way_num - accuracy = 1.0 * correct_predictions.float() / float(total_queries) - - accuracy = torch.tensor([accuracy]) + + accuracy = (predicted_query == ground_truth_query).float().mean() return loss, accuracy @@ -152,8 +140,8 @@ def set_forward_loss(self, batch): query_target, ) = self.split_by_episode(image, mode=2) - loss_list = [] - acc_list = [] + loss_list = torch.zeros(episode_size, device=self.device) + acc_list = torch.zeros(episode_size, device=self.device) for i in range(episode_size): support_label_onehot = self.labels_to_onehot(support_target[i]) @@ -166,13 +154,11 @@ def set_forward_loss(self, batch): query_label_onehot, ) - loss_list.append(loss) - acc_list.append(acc) + loss_list[i] = loss + acc_list[i] = acc - loss = torch.stack(loss_list) - loss = torch.mean(loss) - acc = torch.stack(acc_list) - acc = torch.mean(acc) * 100.0 + loss = loss_list.mean() + acc = acc_list.mean() * 100.0 return None, acc, loss @@ -191,7 +177,7 @@ def set_forward(self, batch): query_target, ) = self.split_by_episode(image, mode=2) - acc_list = [] + acc_list = torch.zeros(episode_size, device=self.device) for i in range(episode_size): support_label_onehot = self.labels_to_onehot(support_target[i]) @@ -204,9 +190,8 @@ def set_forward(self, batch): query_label_onehot, ) - acc_list.append(acc) + acc_list[i] = acc - acc = torch.stack(acc_list) - acc = torch.mean(acc) * 100.0 + acc = acc_list.mean() * 100.0 return None, acc