@@ -13,33 +13,28 @@ class iMAML(GBML):
1313 def __init__ (self , args ):
1414 super ().__init__ (args )
1515 self ._init_net ()
16- self .network_aux = type (self .network )(args ).cuda ()
17- self .network_aux .load_state_dict (self .network .state_dict ())
1816 self ._init_opt ()
19- self .inner_optimizer = torch .optim .SGD (self .network_aux .parameters (), lr = self .args .inner_lr )
2017 self .lamb = 2.0
2118 self .n_cg = 3
2219 return None
2320
2421 @torch .enable_grad ()
25- def inner_loop (self , train_input , train_target ):
26- self . network_aux . zero_grad ()
27- train_logit = self . network_aux (train_input )
22+ def inner_loop (self , fmodel , diffopt , train_input , train_target ):
23+
24+ train_logit = fmodel (train_input )
2825 inner_loss = F .cross_entropy (train_logit , train_target )
29- inner_loss += (self .lamb / 2. ) * ((torch .nn .utils .parameters_to_vector (self .network_aux .parameters ())- torch .nn .utils .parameters_to_vector (self .network .parameters ()).detach ())** 2 ).sum ()
30- inner_loss . backward ( )
31- self . inner_optimizer . step ()
26+ inner_loss += (self .lamb / 2. ) * ((torch .nn .utils .parameters_to_vector (self .network .parameters ())- torch .nn .utils .parameters_to_vector (self .network .parameters ()).detach ())** 2 ).sum ()
27+ diffopt . step ( inner_loss )
28+
3229 return None
3330
34- @torch .enable_grad ()
35- def cg (self , in_grad , outer_grad ):
36- in_grad = torch .nn .utils .parameters_to_vector (in_grad )
37- outer_grad = torch .nn .utils .parameters_to_vector (outer_grad )
31+ @torch .no_grad ()
32+ def cg (self , in_grad , outer_grad , params ):
3833 x = outer_grad .clone ().detach ()
39- r = outer_grad .clone ().detach () - self .hv_prod (in_grad , x )
34+ r = outer_grad .clone ().detach () - self .hv_prod (in_grad , x , params )
4035 p = r .clone ().detach ()
4136 for i in range (self .n_cg ):
42- Ap = self .hv_prod (in_grad , p )
37+ Ap = self .hv_prod (in_grad , p , params )
4338 alpha = (r @ r )/ (p @ Ap )
4439 x = x + alpha * p
4540 r_new = r - alpha * Ap
@@ -59,9 +54,9 @@ def vec_to_grad(self, vec):
5954 pointer += num_param
6055 return res
6156
62- def hv_prod ( self , in_grad , x ):
63- scalar = in_grad @ x . detach ()
64- hv = torch .autograd .grad (scalar , self . network_aux . parameters () , retain_graph = True )
57+ @ torch . enable_grad ()
58+ def hv_prod ( self , in_grad , x , params ):
59+ hv = torch .autograd .grad (in_grad , params , retain_graph = True , grad_outputs = x )
6560 hv = torch .nn .utils .parameters_to_vector (hv ).detach ()
6661 # precondition with identity matrix
6762 return hv / self .lamb + x
@@ -77,27 +72,28 @@ def outer_loop(self, batch, is_train):
7772
7873 for (train_input , train_target , test_input , test_target ) in zip (train_inputs , train_targets , test_inputs , test_targets ):
7974
80- self . network_aux . load_state_dict (self .network . state_dict ())
75+ with higher . innerloop_ctx (self .network , self . inner_optimizer , track_higher_grads = False ) as ( fmodel , diffopt ):
8176
82- for step in range (self .args .n_inner ):
83- self .inner_loop (train_input , train_target )
84-
85- train_logit = self . network_aux (train_input )
86- in_loss = F .cross_entropy (train_logit , train_target )
77+ for step in range (self .args .n_inner ):
78+ self .inner_loop (fmodel , diffopt , train_input , train_target )
79+
80+ train_logit = fmodel (train_input )
81+ in_loss = F .cross_entropy (train_logit , train_target )
8782
88- test_logit = self . network_aux (test_input )
89- outer_loss = F .cross_entropy (test_logit , test_target )
90- loss_log += outer_loss .item ()/ self .batch_size
83+ test_logit = fmodel (test_input )
84+ outer_loss = F .cross_entropy (test_logit , test_target )
85+ loss_log += outer_loss .item ()/ self .batch_size
9186
92- with torch .no_grad ():
93- acc_log += get_accuracy (test_logit , test_target ).item ()/ self .batch_size
94-
95- if is_train :
96- in_grad = torch .autograd .grad (in_loss , self .network_aux .parameters (), create_graph = True )
97- outer_grad = torch .autograd .grad (outer_loss , self .network_aux .parameters ())
98- implicit_grad = self .cg (in_grad , outer_grad )
99- grad_list .append (implicit_grad )
100- loss_list .append (outer_loss .item ())
87+ with torch .no_grad ():
88+ acc_log += get_accuracy (test_logit , test_target ).item ()/ self .batch_size
89+
90+ if is_train :
91+ params = list (fmodel .parameters (time = - 1 ))
92+ in_grad = torch .nn .utils .parameters_to_vector (torch .autograd .grad (in_loss , params , create_graph = True ))
93+ outer_grad = torch .nn .utils .parameters_to_vector (torch .autograd .grad (outer_loss , params ))
94+ implicit_grad = self .cg (in_grad , outer_grad , params )
95+ grad_list .append (implicit_grad )
96+ loss_list .append (outer_loss .item ())
10197
10298 if is_train :
10399 self .outer_optimizer .zero_grad ()
0 commit comments