MCPcopy Create free account
hub / github.com/coperception/star / CoModule

Class CoModule

star/utils/CoModule.py:13–253  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

11
12
13class 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()

Callers 3

mainFunction · 0.90
mainFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by 2

mainFunction · 0.72
mainFunction · 0.72