sub broadcasted cpu autograd
This commit is contained in:
parent
27a866d721
commit
ddf04363d1
12 changed files with 190 additions and 18 deletions
|
|
@ -58,6 +58,45 @@ void sub_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data) {
|
|||
}
|
||||
}
|
||||
|
||||
void sub_broadcasted_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data, int* broadcasted_shape) {
|
||||
int max_ndim = tensor1->ndim > tensor2->ndim ? tensor1->ndim : tensor2->ndim;
|
||||
|
||||
// Calculate strides for broadcasting
|
||||
int* strides1 = (int*)malloc(max_ndim * sizeof(int));
|
||||
int* strides2 = (int*)malloc(max_ndim * sizeof(int));
|
||||
if (strides1 == NULL || strides2 == NULL) {
|
||||
fprintf(stderr, "Memory allocation failed\n");
|
||||
exit(1);
|
||||
}
|
||||
|
||||
int stride1 = 1, stride2 = 1;
|
||||
for (int i = max_ndim - 1; i >= 0; i--) {
|
||||
int dim1 = i < tensor1->ndim ? tensor1->shape[tensor1->ndim - max_ndim + i] : 1;
|
||||
int dim2 = i < tensor2->ndim ? tensor2->shape[tensor2->ndim - max_ndim + i] : 1;
|
||||
strides1[i] = dim1 == broadcasted_shape[i] ? stride1 : 0;
|
||||
strides2[i] = dim2 == broadcasted_shape[i] ? stride2 : 0;
|
||||
stride1 *= broadcasted_shape[i];
|
||||
stride2 *= broadcasted_shape[i];
|
||||
}
|
||||
|
||||
// Perform element-wise addition with broadcasting
|
||||
for (int i = 0; i < tensor1->size; i++) {
|
||||
int index1 = 0, index2 = 0;
|
||||
int linear_index = i;
|
||||
for (int j = max_ndim - 1; j >= 0; j--) {
|
||||
int pos = linear_index % broadcasted_shape[j];
|
||||
linear_index /= broadcasted_shape[j];
|
||||
if (strides1[j] != 0) index1 += pos * strides1[j];
|
||||
if (strides2[j] != 0) index2 += pos * strides2[j];
|
||||
}
|
||||
result_data[i] = tensor1->data[index1] - tensor2->data[index2];
|
||||
}
|
||||
|
||||
// Free strides
|
||||
free(strides1);
|
||||
free(strides2);
|
||||
}
|
||||
|
||||
void elementwise_mul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data) {
|
||||
|
||||
for (int i = 0; i < tensor1->size; i++) {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue