Args: x (float): Input images(BxCxHxW) axis (int): The index for weighted mean other_axis (int): The other index Returns: weighted index for axis, BxC
(x, axis, other_axis, softmax=True)
| 154 | |
| 155 | # |
| 156 | def get_gaussian_mean(x, axis, other_axis, softmax=True): |
| 157 | """ |
| 158 | |
| 159 | Args: |
| 160 | x (float): Input images(BxCxHxW) |
| 161 | axis (int): The index for weighted mean |
| 162 | other_axis (int): The other index |
| 163 | |
| 164 | Returns: weighted index for axis, BxC |
| 165 | |
| 166 | """ |
| 167 | mat2line = torch.sum(x, axis=other_axis) |
| 168 | # mat2line = mat2line / mat2line.mean() * 10 |
| 169 | if softmax: |
| 170 | u = torch.softmax(mat2line, axis=2) |
| 171 | else: |
| 172 | u = mat2line / (mat2line.sum(2, keepdim=True) + 1e-6) |
| 173 | size = x.shape[axis] |
| 174 | ind = torch.linspace(0, 1, size).to(x.device) |
| 175 | batch = x.shape[0] |
| 176 | channel = x.shape[1] |
| 177 | index = ind.repeat([batch, channel, 1]) |
| 178 | mean_position = torch.sum(index * u, dim=2) |
| 179 | return mean_position |
| 180 | |
| 181 | |
| 182 | def get_expected_points_from_map(hm, softmax=True): |
no test coverage detected