mirror of
https://github.com/wassname/baukit.git
synced 2026-07-25 13:00:09 +08:00
Avoid numpy.prod error when computing shape.
This commit is contained in:
@@ -220,7 +220,7 @@ class Stat:
|
||||
setattr(self, attr, data_shape)
|
||||
else:
|
||||
assert x.shape[1:] == data_shape
|
||||
return x.view(x.shape[0], int(numpy.prod(data_shape)))
|
||||
return x.view(x.shape[0], x.shape[1:].numel())
|
||||
|
||||
def _restore_result_shape(self, x, attr="data_shape"):
|
||||
"""
|
||||
@@ -1185,7 +1185,7 @@ class TopK:
|
||||
if self.top_data is None:
|
||||
# Allocation: allocate a buffer of size 5*k, at least 10, for each.
|
||||
self.data_shape = data.shape[1:]
|
||||
feature_size = int(numpy.prod(self.data_shape))
|
||||
feature_size = self.data_shape.numel()
|
||||
self.top_data = torch.zeros(
|
||||
feature_size, max(10, self.k * 5), out=data.new()
|
||||
)
|
||||
@@ -1208,7 +1208,7 @@ class TopK:
|
||||
# Pick: copy the top sk of the next batch into the buffer.
|
||||
# Currently strided topk is slow. So we clone after transpose.
|
||||
# TODO: remove the clone() if it becomes faster.
|
||||
cdata = data.reshape(size, numpy.prod(data.shape[1:])).t().clone()
|
||||
cdata = data.reshape(size, data.shape[1:].numel()).t().clone()
|
||||
td, ti = cdata.topk(sk, sorted=False, largest=self.largest)
|
||||
self.top_data[:, self.next : self.next + sk] = td
|
||||
if index is not None:
|
||||
|
||||
Reference in New Issue
Block a user