1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62
| class MultiScaleDiscriminator(nn.Module): """ 多尺度判别器 同时在多个空间尺度评估生成样本: - 尺度1: 粗略包络 (躯干运动) - 尺度2: 中等纹理 (肢体摆动) - 尺度3: 精细细节 (微振动) """ def __init__(self, spec_channels=1, base_channels=64): super().__init__() self.disc_coarse = self._make_disc_block(spec_channels, base_channels, kernel_size=7, stride=2) self.disc_medium = self._make_disc_block(spec_channels, base_channels, kernel_size=5, stride=1) self.disc_fine = self._make_disc_block(spec_channels, base_channels, kernel_size=3, stride=1) self.fusion = nn.Sequential( nn.Conv2d(base_channels * 3, base_channels, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_channels, 1, 4, padding=0), nn.Sigmoid() ) def _make_disc_block(self, in_ch, out_ch, kernel_size, stride): return nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size, stride=stride, padding=kernel_size//2), nn.GroupNorm(8, out_ch), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(out_ch, out_ch, kernel_size, stride=2, padding=kernel_size//2), nn.GroupNorm(8, out_ch), nn.LeakyReLU(0.2, inplace=True) ) def forward(self, spec): feat_coarse = self.disc_coarse(spec) feat_medium = self.disc_medium(spec) feat_fine = self.disc_fine(spec) target_size = feat_coarse.shape[2:] feat_medium = nn.functional.interpolate(feat_medium, size=target_size) feat_fine = nn.functional.interpolate(feat_fine, size=target_size) combined = torch.cat([feat_coarse, feat_medium, feat_fine], dim=1) output = self.fusion(combined) return output, [feat_coarse, feat_medium, feat_fine]
disc = MultiScaleDiscriminator() validity, features = disc(fake_spec.detach()) print(f"判别器输出: {validity.shape}") print(f"多尺度特征: {[f.shape for f in features]}")
|