| 11 | |
| 12 | |
| 13 | class CoModule(object): |
| 14 | def __init__(self, model, optimizer, com): |
| 15 | self.mae_loss_scaler = NativeScaler() |
| 16 | self.model = model |
| 17 | self.optimizer = optimizer |
| 18 | self.scheduler = None |
| 19 | if com=="late" or com=="vqvae": |
| 20 | self.scheduler = torch.optim.lr_scheduler.MultiStepLR( |
| 21 | optimizer, milestones=[50, 100, 150, 200], gamma=0.5 |
| 22 | ) |
| 23 | |
| 24 | def resume_from_cpu(self, checkpoint, device, trainable=True): |
| 25 | """ |
| 26 | This function load state dict to model and optimizer on cpu, and move it back to device. |
| 27 | This avoids a GPU memory surge issue. |
| 28 | NOTE: assume checkpoint is loaded in cpu |
| 29 | """ |
| 30 | # handles model |
| 31 | self.model = self.model.cpu() |
| 32 | self.model.load_state_dict(checkpoint["model_state_dict"]) |
| 33 | self.model = self.model.to(device) |
| 34 | if trainable: |
| 35 | # handles optimizer |
| 36 | self.optimizer.load_state_dict(checkpoint["optimizer_state_dict"]) |
| 37 | optimizer_to(self.optimizer, device) |
| 38 | # possible extension: reinitialize scheduler based on this new optimizer |
| 39 | self.scheduler = self.scheduler = torch.optim.lr_scheduler.MultiStepLR( |
| 40 | self.optimizer, milestones=[50, 100, 150, 200], gamma=0.5 |
| 41 | ) |
| 42 | |
| 43 | # used by scene completion task |
| 44 | def step_completion(self, data, batch_size, loss_fn='ce', trainable=False): |
| 45 | bev_seq = data['bev_seq'] |
| 46 | trans_matrices = data['trans_matrices'] |
| 47 | num_agent = data['num_agent'] |
| 48 | |
| 49 | result, ind_pred = self.model(bev_seq, trans_matrices, num_agent, batch_size=batch_size) |
| 50 | |
| 51 | loss_fn_dict = { |
| 52 | 'mse': nn.MSELoss(), |
| 53 | 'bce': nn.BCELoss(), |
| 54 | 'ce': nn.CrossEntropyLoss(), |
| 55 | 'l1': nn.L1Loss(), |
| 56 | 'smooth_l1': nn.SmoothL1Loss(), |
| 57 | } |
| 58 | |
| 59 | loss = -1 |
| 60 | if trainable: |
| 61 | # labels = data['bev_seq_teacher'] |
| 62 | # labels = labels.permute(0, 1, 4, 2, 3).squeeze() # (Batch, seq, z, h, w) |
| 63 | # loss = 10000 * loss_fn_dict[loss_fn](result, labels) |
| 64 | target = bev_seq.permute(0, 1, 4, 2, 3).squeeze(1) |
| 65 | target = target.type(torch.LongTensor).to(ind_pred.device) |
| 66 | loss = loss_fn_dict[loss_fn](ind_pred, target) |
| 67 | |
| 68 | if self.MGDA: |
| 69 | self.optimizer_encoder.zero_grad() |
| 70 | self.optimizer_head.zero_grad() |