MCPcopy Create free account
hub / github.com/IntelligentSystemsLab/ST-EVCDP / physics_informed_meta_learning

Function physics_informed_meta_learning

learner.py:10–80  ·  view source on GitHub ↗
(law_list, global_model, model_name, p_epoch, bs, train_occupancy, train_price, seq_l, pre_l, device, adj_dense)

Source from the content-addressed store, hash-verified

8
9#
10def 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()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected