Skip to content

Commit 54aa9c6

Browse files
Precompute reverse dim mapping for tensor slicing (closes #230)
1 parent f4d122d commit 54aa9c6

2 files changed

Lines changed: 52 additions & 35 deletions

File tree

src/interpreter.c

Lines changed: 26 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -3078,33 +3078,41 @@ ExecResult assign_index_chain(Interpreter *interp, Env *env, Expr *idx_expr, Val
30783078
goto cleanup;
30793079
}
30803080

3081+
int *out_to_orig = malloc(sizeof(int) * new_ndim);
3082+
if (!out_to_orig) {
3083+
free(starts);
3084+
free(ends);
3085+
free(orig_to_out);
3086+
free(out_shape);
3087+
out = make_error("Out of memory", stmt_line, stmt_col);
3088+
goto cleanup;
3089+
}
3090+
for (size_t i = 0; i < t->ndim; i++) {
3091+
if (orig_to_out[i] >= 0) {
3092+
out_to_orig[orig_to_out[i]] = (int)i;
3093+
}
3094+
}
3095+
3096+
size_t fixed_dim_offset = 0;
3097+
for (size_t k = 0; k < t->ndim; k++) {
3098+
if (orig_to_out[k] == -1) {
3099+
size_t pos = (ends[k] >= starts[k]) ? (size_t)(starts[k] - 1) : 0;
3100+
fixed_dim_offset += pos * t->strides[k];
3101+
}
3102+
}
3103+
30813104
// Write RHS elements into target tensor region
30823105
// Iterate over output positions and compute corresponding source offset
30833106
for (size_t out_idx = 0; out_idx < rt->length; out_idx++) {
3084-
// compute multi-index for out
30853107
size_t rem = out_idx;
3086-
size_t src_offset = 0;
3108+
size_t src_offset = fixed_dim_offset;
30873109
for (size_t d = 0; d < new_ndim; d++) {
30883110
size_t pos = rem / rt->strides[d];
30893111
rem = rem % rt->strides[d];
3090-
// find orig dim for this d
3091-
size_t orig = 0;
3092-
for (size_t k = 0; k < t->ndim; k++) {
3093-
if (orig_to_out[k] == (int)d) {
3094-
orig = k;
3095-
break;
3096-
}
3097-
}
3112+
size_t orig = (size_t)out_to_orig[d];
30983113
size_t src_pos = pos + (size_t)(starts[orig] - 1);
30993114
src_offset += src_pos * t->strides[orig];
31003115
}
3101-
// add fixed-dimension offsets
3102-
for (size_t k = 0; k < t->ndim; k++) {
3103-
if (orig_to_out[k] == -1) {
3104-
size_t pos = (ends[k] >= starts[k]) ? (size_t)(starts[k] - 1) : 0;
3105-
src_offset += pos * t->strides[k];
3106-
}
3107-
}
31083116

31093117
// assign element
31103118
mtx_lock(&t->lock);
@@ -3117,6 +3125,7 @@ ExecResult assign_index_chain(Interpreter *interp, Env *env, Expr *idx_expr, Val
31173125
free(starts);
31183126
free(ends);
31193127
free(orig_to_out);
3128+
free(out_to_orig);
31203129
// After slice assignment, set cur to base (no further chaining into this node)
31213130
cur = &base_val;
31223131
rhs_applied_directly = true;

src/value.c

Lines changed: 26 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -371,43 +371,51 @@ Value value_tensor_slice(Value v, const int64_t *starts, const int64_t *ends, si
371371
}
372372
}
373373

374+
int *out_to_orig = malloc(sizeof(int) * new_ndim);
375+
if (!out_to_orig) {
376+
free(new_shape);
377+
free(nstarts);
378+
free(nends);
379+
free(orig_to_out);
380+
fprintf(stderr, "Out of memory\n");
381+
exit(1);
382+
}
383+
for (size_t i = 0; i < t->ndim; i++) {
384+
if (orig_to_out[i] >= 0) {
385+
out_to_orig[orig_to_out[i]] = (int)i;
386+
}
387+
}
388+
389+
size_t fixed_dim_offset = 0;
390+
for (size_t k = 0; k < t->ndim; k++) {
391+
if (orig_to_out[k] == -1) {
392+
size_t pos = (nends[k] >= nstarts[k]) ? (size_t)(nstarts[k] - 1) : 0;
393+
fixed_dim_offset += pos * t->strides[k];
394+
}
395+
}
396+
374397
Value out = value_tensor_new(t->elem_type, new_ndim, new_shape);
375398
Tensor *ot = out.as.tensor;
376399

377400
// iterate over output positions and copy corresponding element
378401
for (size_t out_idx = 0; out_idx < ot->length; out_idx++) {
379-
// compute multi-index for out
380402
size_t rem = out_idx;
381-
size_t src_offset = 0;
403+
size_t src_offset = fixed_dim_offset;
382404
for (size_t d = 0; d < new_ndim; d++) {
383405
size_t pos = rem / ot->strides[d];
384406
rem = rem % ot->strides[d];
385-
// find corresponding original dimension
386-
// scan orig_to_out to find index with value == d
387-
size_t orig = 0;
388-
for (size_t k = 0; k < t->ndim; k++) {
389-
if (orig_to_out[k] == (int)d) {
390-
orig = k;
391-
break;
392-
}
393-
}
407+
size_t orig = (size_t)out_to_orig[d];
394408
size_t src_pos = pos + (size_t)(nstarts[orig] - 1);
395409
src_offset += src_pos * t->strides[orig];
396410
}
397-
// add fixed-dimension offsets
398-
for (size_t k = 0; k < t->ndim; k++) {
399-
if (orig_to_out[k] == -1) {
400-
size_t pos = (nends[k] >= nstarts[k]) ? (size_t)(nstarts[k] - 1) : 0;
401-
src_offset += pos * t->strides[k];
402-
}
403-
}
404411
ot->data[out_idx] = value_copy(t->data[src_offset]);
405412
}
406413

407414
free(new_shape);
408415
free(nstarts);
409416
free(nends);
410417
free(orig_to_out);
418+
free(out_to_orig);
411419
return out;
412420
}
413421

0 commit comments

Comments
 (0)