(law_list, model, model_name, p_epoch, bs, train_occupancy, train_price, seq_l, pre_l, device, adj_dense)
| 81 | |
| 82 | |
| 83 | def fast_learning(law_list, model, model_name, p_epoch, bs, train_occupancy, train_price, seq_l, pre_l, device, adj_dense): |
| 84 | n_laws = len(law_list) |
| 85 | fast_datasets = dict() |
| 86 | fast_loaders = dict() |
| 87 | for n in range(n_laws): |
| 88 | fast_datasets[n] = fn.CreateFastDataset(train_occupancy, train_price, seq_l, pre_l, law_list[n], device, adj_dense) |
| 89 | fast_loaders[n] = DataLoader(fast_datasets[n], batch_size=bs, shuffle=True, drop_last=True) |
| 90 | |
| 91 | optimizer = torch.optim.Adam(model.parameters(), weight_decay=0.00001) |
| 92 | loss_function = torch.nn.MSELoss() |
| 93 | for epoch in tqdm(range(p_epoch), desc='Pre-training'): |
| 94 | for n in range(n_laws): |
| 95 | for j, data in enumerate(fast_loaders[n]): |
| 96 | ''' |
| 97 | occupancy = (batch, seq, node) |
| 98 | price = (batch, seq, node) |
| 99 | label = (batch, node) |
| 100 | ''' |
| 101 | occupancy, price, label, prc_ch, label_ch = data |
| 102 | optimizer.zero_grad() |
| 103 | predict = model(occupancy, prc_ch) |
| 104 | loss = loss_function(predict, label_ch) |
| 105 | loss.backward() |
| 106 | optimizer.step() |
| 107 | |
| 108 | for j, data in enumerate(fast_loaders[n]): |
| 109 | ''' |
| 110 | occupancy = (batch, seq, node) |
| 111 | price = (batch, seq, node) |
| 112 | label = (batch, node) |
| 113 | ''' |
| 114 | occupancy, price, label, prc_ch, label_ch = data |
| 115 | optimizer.zero_grad() |
| 116 | predict = model(occupancy, prc_ch) |
| 117 | loss = loss_function(predict, label_ch) |
| 118 | loss.backward() |
| 119 | optimizer.step() |
| 120 | |
| 121 | return model |
nothing calls this directly
no outgoing calls
no test coverage detected