update to support torch 2.0

This commit is contained in:
liulixinkerry 2023-05-27 15:54:45 +08:00
parent 9cc7ea4a97
commit 0aa58ed4a5
2 changed files with 5 additions and 6 deletions

View File

@ -11,7 +11,7 @@ def get_option():
parser.add_argument('--load_from_raw', type=str2bool, default=True, help='If True, parse and load from benchmark files. If False, load from pt') parser.add_argument('--load_from_raw', type=str2bool, default=True, help='If True, parse and load from benchmark files. If False, load from pt')
parser.add_argument('--run_all', type=str2bool, default=False, help='If True, run all designs in the given dataset. If False, run the given design_name only.') parser.add_argument('--run_all', type=str2bool, default=False, help='If True, run all designs in the given dataset. If False, run the given design_name only.')
parser.add_argument('--seed', type=int, default=0, help='seed to initialize all the random modules') parser.add_argument('--seed', type=int, default=0, help='seed to initialize all the random modules')
parser.add_argument('--gpu', type=int, default=1, help='gpu id') parser.add_argument('--gpu', type=int, default=0, help='gpu id')
parser.add_argument('--num_threads', type=int, default=20, help='threads') parser.add_argument('--num_threads', type=int, default=20, help='threads')
parser.add_argument('--deterministic', type=str2bool, default=True, help='use deterministic mode') parser.add_argument('--deterministic', type=str2bool, default=True, help='use deterministic mode')

View File

@ -19,7 +19,7 @@ def apply_precond(mov_node_pos: torch.Tensor, ps: ParamScheduler, args):
mov_node_pos.grad /= ps.precond_weight mov_node_pos.grad /= ps.precond_weight
return mov_node_pos.grad return mov_node_pos.grad
# For nesterov
def calc_obj_and_grad( def calc_obj_and_grad(
mov_node_pos, mov_node_pos,
constraint_fn=None, constraint_fn=None,
@ -103,18 +103,17 @@ def calc_obj_and_grad(
grad = apply_precond(mov_node_pos, ps, args) grad = apply_precond(mov_node_pos, ps, args)
return loss, grad return loss, grad
# For Adam
def calc_grad( def calc_grad(
optimizer: torch.optim.Optimizer, mov_node_pos: torch.Tensor, wl_loss, density_loss optimizer: torch.optim.Optimizer, mov_node_pos: torch.Tensor, wl_loss, density_loss
): ):
optimizer.zero_grad() optimizer.zero_grad(set_to_none=False)
wl_loss.backward(retain_graph=True) wl_loss.backward(retain_graph=True)
wl_grad = mov_node_pos.grad.detach().clone() wl_grad = mov_node_pos.grad.detach().clone()
optimizer.zero_grad() optimizer.zero_grad(set_to_none=False)
density_loss.backward(retain_graph=True) density_loss.backward(retain_graph=True)
density_grad = mov_node_pos.grad.detach().clone() density_grad = mov_node_pos.grad.detach().clone()
optimizer.zero_grad() optimizer.zero_grad(set_to_none=False)
return wl_grad, density_grad return wl_grad, density_grad