fix add min max keepdim
This commit is contained in:
parent
5e72fe57d3
commit
fd831f950d
4 changed files with 39 additions and 45 deletions
BIN
build/tensor.o
BIN
build/tensor.o
Binary file not shown.
|
|
@ -183,11 +183,8 @@ extern "C" {
|
|||
int ndim;
|
||||
int* shape;
|
||||
if (axis == -1) {
|
||||
|
||||
shape = (int*) malloc(sizeof(int));
|
||||
if (shape == NULL) {
|
||||
fprintf(stderr, "Memory allocation failed\n");
|
||||
exit(1);
|
||||
}
|
||||
shape[0] = 1;
|
||||
ndim = 1;
|
||||
} else {
|
||||
|
|
@ -198,7 +195,6 @@ extern "C" {
|
|||
}
|
||||
}
|
||||
ndim = tensor->ndim - 1;
|
||||
|
||||
}
|
||||
|
||||
int size = 1;
|
||||
|
|
@ -220,25 +216,27 @@ extern "C" {
|
|||
fprintf(stderr, "Memory allocation failed\n");
|
||||
exit(1);
|
||||
}
|
||||
|
||||
sum_tensor_cpu(tensor, result_data, size, shape, axis);
|
||||
|
||||
if (keepdim) {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
if (axis == -1) {
|
||||
if (axis == -1){
|
||||
ndim = tensor->ndim;
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = 1;
|
||||
} else {
|
||||
}
|
||||
} else {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = tensor->shape[i];
|
||||
}
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
}
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
Tensor* new_tensor = create_tensor(result_data, shape, ndim, device);
|
||||
return reshape_tensor(new_tensor, shape, ndim);
|
||||
|
||||
}
|
||||
|
||||
return create_tensor(result_data, shape, ndim, device);
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -252,12 +250,8 @@ extern "C" {
|
|||
}
|
||||
int ndim;
|
||||
int* shape;
|
||||
if (axis == -1) {
|
||||
if (axis == -1) {
|
||||
shape = (int*) malloc(sizeof(int));
|
||||
if (shape == NULL) {
|
||||
fprintf(stderr, "Memory allocation failed\n");
|
||||
exit(1);
|
||||
}
|
||||
shape[0] = 1;
|
||||
ndim = 1;
|
||||
} else {
|
||||
|
|
@ -268,7 +262,6 @@ extern "C" {
|
|||
}
|
||||
}
|
||||
ndim = tensor->ndim - 1;
|
||||
|
||||
}
|
||||
|
||||
int size = 1;
|
||||
|
|
@ -293,18 +286,20 @@ extern "C" {
|
|||
max_tensor_cpu(tensor, result_data, size, shape, axis);
|
||||
|
||||
if (keepdim) {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
if (axis == -1) {
|
||||
if (axis == -1){
|
||||
ndim = tensor->ndim;
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = 1;
|
||||
} else {
|
||||
}
|
||||
} else {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = tensor->shape[i];
|
||||
}
|
||||
}
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
Tensor* new_tensor = create_tensor(result_data, shape, ndim, device);
|
||||
return reshape_tensor(new_tensor, shape, ndim);
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
}
|
||||
}
|
||||
|
||||
return create_tensor(result_data, shape, ndim, device);
|
||||
|
|
@ -322,11 +317,8 @@ extern "C" {
|
|||
int ndim;
|
||||
int* shape;
|
||||
if (axis == -1) {
|
||||
|
||||
shape = (int*) malloc(sizeof(int));
|
||||
if (shape == NULL) {
|
||||
fprintf(stderr, "Memory allocation failed\n");
|
||||
exit(1);
|
||||
}
|
||||
shape[0] = 1;
|
||||
ndim = 1;
|
||||
} else {
|
||||
|
|
@ -337,9 +329,8 @@ extern "C" {
|
|||
}
|
||||
}
|
||||
ndim = tensor->ndim - 1;
|
||||
|
||||
}
|
||||
|
||||
|
||||
int size = 1;
|
||||
for (int i = 0; i < ndim; i++) {
|
||||
size *= shape[i];
|
||||
|
|
@ -362,18 +353,20 @@ extern "C" {
|
|||
min_tensor_cpu(tensor, result_data, size, shape, axis);
|
||||
|
||||
if (keepdim) {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
if (axis == -1) {
|
||||
if (axis == -1){
|
||||
ndim = tensor->ndim;
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = 1;
|
||||
} else {
|
||||
}
|
||||
} else {
|
||||
shape = (int*) malloc((tensor->ndim) * sizeof(int));
|
||||
for (int i = 0; i < tensor->ndim; i++) {
|
||||
shape[i] = tensor->shape[i];
|
||||
}
|
||||
}
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
Tensor* new_tensor = create_tensor(result_data, shape, ndim, device);
|
||||
return reshape_tensor(new_tensor, shape, ndim);
|
||||
shape[axis] = 1;
|
||||
ndim = tensor->ndim;
|
||||
}
|
||||
}
|
||||
|
||||
return create_tensor(result_data, shape, ndim, device);
|
||||
|
|
|
|||
Binary file not shown.
|
|
@ -272,6 +272,7 @@ class TestTensorOperations(unittest.TestCase):
|
|||
torch_tensor = torch.tensor([[[1, 2], [3, -4]], [[5, 6], [7, 8]]]).to(self.device)
|
||||
torch_expected = torch.sum(torch_tensor, dim=1, keepdim=True)
|
||||
|
||||
print(torch_result, torch_expected)
|
||||
self.assertTrue(utils.compare_torch(torch_result, torch_expected))
|
||||
|
||||
def test_max(self):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue