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

Function get_gaussian_mean

util/utils.py:156–179  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

154
155#
156def 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
182def get_expected_points_from_map(hm, softmax=True):

Callers 1

Calls 1

toMethod · 0.45

Tested by

no test coverage detected