MCPcopy Create free account
hub / github.com/OpenGVLab/HumanBench / add_aio_decoder_specific

Function add_aio_decoder_specific

PATH/core/utils.py:626–661  ·  view source on GitHub ↗
(m, decoder_specific, task_sp_list=(), neck_sp_list=())

Source from the content-addressed store, hash-verified

624 printlog('add buffer {} as neck_specific'.format(name))
625
626def 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
664def copy_state_dict_cpu(state_dict):

Callers 1

__init__Method · 0.90

Calls 1

printlogFunction · 0.85

Tested by

no test coverage detected