o `sz*MultiBand_MoE.__init__..N) r r3rrbackbonerr patch_embedproj ModuleListrangemulti_band_experts)rr6r7rErrrrQs zMultiBand_MoE.__init__c Cs"t|dt|jddf|j}tt|jD],}|j||dd||dddddf}|d}||dd|ddddf<qt j |dd}||}|j |}|ddddddf}|j \}}} t|d} } ||| | | }|dddd}||ddddddffS)u 前向传播函数。 参数: - input_data: 输入数据,形状为 (batch_size, num_bands, h, w)。 - psf: 点扩散函数,形状为 (batch_size, num_bands, h, w)。 返回: - output: 特征图像,形状为 (batch_size, num_bands, h, w)。 - multi_band_weights: 多波段权重,形状为 (batch_size, num_bands, 1, 1)。 rrNr'g?r)torchzerossizelenrEtodevicerD unsqueezer/softmaxr@forward_featuresshapeintreshapepermute) rrZmulti_band_weightsiZmulti_band_weightZ output_datafeatureBNCHWrrrrbs& .   zMultiBand_MoE.forward)r4Fr5r rrrrr3Osr3cr) ReconTask_MoEcstt|t||||_||_td||_t tj |jdddddt dt tj ddddddt dt tj ddddddt dt tj dd dddd |_ dS) Nr9@r rFr)rr;r r:r5)r r[rr%gatingr.rQ embeding_dimrrConvTranspose2drrdecoderr,rrrrs  zReconTask_MoE.__init__cshfddtjDdj\}}}tfddtjD}|}|fS)Ncs<g|]}dd|j|djddddfqSNr)r_r=rTr1rrr?s<z)ReconTask_MoE.forward..rc3s4|]}dd|fddd|VqdSrb)viewrc)bexpert_outputs gate_weightsrr s2z(ReconTask_MoE.forward..)r^rDr.rPsumra)rr2chwoutputr)rerfrgrr2rrs   zReconTask_MoE.forwardr rrrrr[s r[) rGtorch.nnrZtorch.nn.functional functionalr/Z timm.modelsrModulerr%r3r[rrrrs  -: