(self,dataloader,class_list,session)
| 56 | raise ValueError('Unknown mode') |
| 57 | |
| 58 | def update_fc(self,dataloader,class_list,session): |
| 59 | for batch in dataloader: |
| 60 | data, label = [_.cuda() for _ in batch] |
| 61 | data=self.encode(data).detach() |
| 62 | |
| 63 | if self.args.not_data_init: |
| 64 | new_fc = nn.Parameter( |
| 65 | torch.rand(len(class_list), self.num_features, device="cuda"), |
| 66 | requires_grad=True) |
| 67 | nn.init.kaiming_uniform_(new_fc, a=math.sqrt(5)) |
| 68 | else: |
| 69 | new_fc = self.update_fc_avg(data, label, class_list) |
| 70 | |
| 71 | if 'ft' in self.args.new_mode: # further finetune |
| 72 | self.update_fc_ft(new_fc,data,label,session) |
| 73 | |
| 74 | def update_fc_avg(self,data,label,class_list): |
| 75 | new_fc=[] |
no test coverage detected