(law_list, global_model, model_name, p_epoch, bs, train_occupancy, train_price, seq_l, pre_l, device, adj_dense)
| 8 | |
| 9 | # |
| 10 | def physics_informed_meta_learning(law_list, global_model, model_name, p_epoch, bs, train_occupancy, train_price, seq_l, pre_l, device, adj_dense): |
| 11 | support_occ, query_occ = fn.meta_division(train_occupancy, support_rate=0.5, query_rate=0.5) |
| 12 | support_prc, query_prc = fn.meta_division(train_price, support_rate=0.5, query_rate=0.5) |
| 13 | |
| 14 | # pre-training data generation |
| 15 | n_laws = len(law_list) |
| 16 | support_dataset_dict = dict() |
| 17 | query_dataset_dict = dict() |
| 18 | support_dataloader_dict = dict() |
| 19 | query_dataloader_dict = dict() |
| 20 | for n in range(n_laws): |
| 21 | support_dataset_dict[n] = fn.PseudoDataset(support_occ, support_prc, seq_l, pre_l, device, adj_dense, law_list[n]) |
| 22 | query_dataset_dict[n] = fn.PseudoDataset(query_occ, query_prc, seq_l, pre_l, device, adj_dense, law_list[n]) |
| 23 | support_dataloader_dict[n] = DataLoader(support_dataset_dict[n], batch_size=bs, shuffle=True, drop_last=True) |
| 24 | query_dataloader_dict[n] = DataLoader(query_dataset_dict[n], batch_size=query_occ.shape[0], shuffle=False) |
| 25 | |
| 26 | # meta-learning process |
| 27 | torch.save(global_model, './checkpoints' + '/meta_' + model_name + '_' + str(pre_l) + '_bs' + str(bs) + 'model.pt') |
| 28 | loss_function = torch.nn.MSELoss() |
| 29 | # outer loop |
| 30 | global_model.train() |
| 31 | for epoch in tqdm(range(p_epoch), desc='Pre-training'): |
| 32 | query_loss = 100 |
| 33 | global_grads = fn.zero_init_global_gradient(global_model) |
| 34 | |
| 35 | # inner loop |
| 36 | for n in range(n_laws): |
| 37 | temp_model = torch.load('./checkpoints' + '/meta_' + model_name + '_' + str(pre_l) + '_bs' + str(bs) + 'model.pt').to(device) |
| 38 | temp_optimizer = torch.optim.Adam(temp_model.parameters(), weight_decay=0.00001) |
| 39 | temp_model.train() |
| 40 | # support |
| 41 | for j, data in enumerate(support_dataloader_dict[n]): |
| 42 | ''' |
| 43 | occupancy = (batch, seq, node) |
| 44 | price = (batch, seq, node) |
| 45 | label = (batch, node) |
| 46 | ''' |
| 47 | occupancy, price, label, pseudo_price, pseudo_label = data |
| 48 | mix_ratio = (j+1) * occupancy.shape[0] / len(train_occupancy) |
| 49 | mix_prc = fn.data_mix(price, pseudo_price, mix_ratio) |
| 50 | mix_label = fn.data_mix(label, pseudo_label, mix_ratio) |
| 51 | temp_optimizer.zero_grad() |
| 52 | predict = temp_model(occupancy, mix_prc) |
| 53 | loss = loss_function(predict, mix_label) |
| 54 | loss.backward() |
| 55 | temp_optimizer.step() |
| 56 | # query |
| 57 | for j, data in enumerate(query_dataloader_dict[n]): |
| 58 | ''' |
| 59 | occupancy = (batch, seq, node) |
| 60 | price = (batch, seq, node) |
| 61 | label = (batch, node) |
| 62 | ''' |
| 63 | occupancy, price, label, pseudo_price, pseudo_label = data |
| 64 | temp_optimizer.zero_grad() |
| 65 | predict = temp_model(occupancy, price) |
| 66 | loss = loss_function(predict, label) |
| 67 | loss.backward() |
nothing calls this directly
no outgoing calls
no test coverage detected