MCPcopy Create free account
hub / github.com/MotrixLab/AiOS / dn_post_process2

Method dn_post_process2

models/aios/aios_smplx.py:2940–2964  ·  view source on GitHub ↗
(self, outputs_class, outputs_coord, mask_dict)

Source from the content-addressed store, hash-verified

2938 return input_query_label, input_query_bbox, attn_mask, attn_mask2, attn_mask3, mask_dict
2939
2940 def dn_post_process2(self, outputs_class, outputs_coord, mask_dict):
2941 if mask_dict and mask_dict['pad_size'] > 0:
2942 output_known_class = [
2943 outputs_class_i[:, :mask_dict['pad_size'], :]
2944 for outputs_class_i in outputs_class
2945 ]
2946 output_known_coord = [
2947 outputs_coord_i[:, :mask_dict['pad_size'], :]
2948 for outputs_coord_i in outputs_coord
2949 ]
2950
2951 outputs_class = [
2952 outputs_class_i[:, mask_dict['pad_size']:, :]
2953 for outputs_class_i in outputs_class
2954 ]
2955 outputs_coord = [
2956 outputs_coord_i[:, mask_dict['pad_size']:, :]
2957 for outputs_coord_i in outputs_coord
2958 ]
2959
2960 mask_dict.update({
2961 'output_known_coord': output_known_coord,
2962 'output_known_class': output_known_class
2963 })
2964 return outputs_class, outputs_coord
2965
2966 def forward(self, data_batch: NestedTensor, targets: List = None):
2967 """The forward expects a NestedTensor, which consists of:

Callers 1

forwardMethod · 0.95

Calls 1

updateMethod · 0.45

Tested by

no test coverage detected