| 624 | printlog('add buffer {} as neck_specific'.format(name)) |
| 625 | |
| 626 | def add_aio_decoder_specific(m, decoder_specific, task_sp_list=(), neck_sp_list=()): |
| 627 | for name, param in m.named_parameters(): |
| 628 | _task_sp_flag = any(name.startswith(sp_name) or name.endswith(sp_name) for sp_name in task_sp_list) |
| 629 | _neck_sp_flag = any(name.startswith(sp_name) or name.endswith(sp_name) for sp_name in neck_sp_list) |
| 630 | |
| 631 | param.task_specific = _task_sp_flag |
| 632 | param.backbone_specific = False |
| 633 | param.neck_specific = _neck_sp_flag |
| 634 | param.decoder_specific = False if _task_sp_flag or _neck_sp_flag else decoder_specific |
| 635 | |
| 636 | if _task_sp_flag: |
| 637 | printlog('add param {} as task_specific'.format(name)) |
| 638 | elif _neck_sp_flag: |
| 639 | printlog('add param {} as neck_specific'.format(name)) |
| 640 | elif decoder_specific: |
| 641 | printlog('add param {} as decoder_specific'.format(name)) |
| 642 | |
| 643 | if not hasattr(torch.nn.Module, 'named_buffers'): |
| 644 | printlog('registering named_buffers for nn.Module at add_decoder_specific') |
| 645 | torch.nn.Module.named_buffers = named_buffers |
| 646 | |
| 647 | #m.cuda() # neccesary for broadcast in DistModule, since buffers are tensors which will be changed after .cuda() |
| 648 | for name, buffer in m.named_buffers(): |
| 649 | _task_sp_flag = any(name.startswith(sp_name) or name.endswith(sp_name) for sp_name in task_sp_list) |
| 650 | _neck_sp_flag = any(name.startswith(sp_name) or name.endswith(sp_name) for sp_name in neck_sp_list) |
| 651 | |
| 652 | buffer.task_specific = _task_sp_flag |
| 653 | buffer.backbone_specific = False |
| 654 | buffer.neck_specific = _neck_sp_flag |
| 655 | buffer.decoder_specific = False if _task_sp_flag or _neck_sp_flag else decoder_specific |
| 656 | if _task_sp_flag: |
| 657 | printlog('add buffer {} as task_specific'.format(name)) |
| 658 | elif _neck_sp_flag: |
| 659 | printlog('add buffer {} as neck_specific'.format(name)) |
| 660 | elif decoder_specific: |
| 661 | printlog('add buffer {} as decoder_specific'.format(name)) |
| 662 | |
| 663 | |
| 664 | def copy_state_dict_cpu(state_dict): |