Skip to content

Commit ddcea3d

Browse files
Transurgeonclaudedance858
authored
Adds indexing atom implementation and other bindings (#20)
* Add index atom for array indexing and slicing Implements efficient indexing operator that handles both CVXPY's `index` (slice-based) and `special_index` (array/boolean) atoms. - Forward: O(n_selected) gather operation - Jacobian: Pre-computed row mapping for fast memcpy per row - Hessian: Pre-allocated scatter buffer with accumulation for repeated indices (equivalent to np.add.at) Enables NLP tests like test_hs071 and test_rosenbrock that use x[0], x[1] style indexing. Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Add reshape converter (Fortran order only) Reshape with order='F' is a pass-through since the underlying data layout is unchanged - only shape interpretation differs. Note: Only Fortran order is supported. C order would require data permutation which is not yet implemented. Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Add quad_over_lin binding and converter Enables test_socp and test_portfolio_socp to pass. Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Add rel_entr binding and converter Exposes the existing rel_entr C implementation (for equal-sized args) through Python bindings. The scalar argument variants declared in the header are not yet implemented in C. Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Optimize index atom: add duplicate detection, simplify Jacobian - Add has_duplicates flag to detect repeated indices at construction - Hessian eval: fast path (direct write) when no duplicates, avoiding O(child_size) memset; slow path (zero + accumulate) for duplicates - Remove jac_row_starts and jac_row_lengths arrays from index_expr - Simplify jacobian_init: use CSR p array directly instead of helpers - Simplify eval_jacobian: compute row lengths on the fly (trivial cost) - Convert &array[i] to array + i style for consistency with codebase - Update CLAUDE.md with improved build instructions and atom lists Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * small edits, most substantial thing removed unnecessary loop --------- Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com> Co-authored-by: Daniel <danielcederberg1@gmail.com>
1 parent 749420d commit ddcea3d

12 files changed

Lines changed: 673 additions & 25 deletions

File tree

‎CLAUDE.md‎

Lines changed: 26 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -11,17 +11,26 @@ DNLP-diff-engine is a C library with Python bindings that provides automatic dif
1111
### Python Package (Recommended)
1212

1313
```bash
14-
# Install in development mode (builds C library + Python bindings)
15-
pip install -e .
14+
# Install in development mode with uv (recommended)
15+
uv pip install -e ".[test]"
1616

17-
# Run all Python tests
17+
# Or with pip
18+
pip install -e ".[test]"
19+
20+
# Run all Python tests (tests are in python/tests/)
1821
pytest
1922

2023
# Run specific test file
21-
pytest tests/python/test_unconstrained.py
24+
pytest python/tests/test_unconstrained.py
2225

2326
# Run specific test
24-
pytest tests/python/test_unconstrained.py::test_sum_log
27+
pytest python/tests/test_unconstrained.py::test_sum_log
28+
29+
# Lint with ruff
30+
ruff check src/
31+
32+
# Auto-fix lint issues
33+
ruff check --fix src/
2534
```
2635

2736
### Standalone C Library
@@ -35,16 +44,6 @@ cmake --build build
3544
./build/all_tests
3645
```
3746

38-
### Legacy Python Build (without pip)
39-
40-
```bash
41-
# Build Python bindings manually (from project root)
42-
cd python && cmake -B build -S . && cmake --build build && cd ..
43-
44-
# Run tests from python/ directory
45-
cd python && python tests/test_problem_native.py
46-
```
47-
4847
## Architecture
4948

5049
### Expression Tree System
@@ -59,10 +58,10 @@ The core abstraction is the `expr` struct (in `include/expr.h`) representing a n
5958

6059
Atoms are organized by mathematical properties in `src/`:
6160

62-
- **`affine/`** - Linear operations: `variable`, `constant`, `add`, `neg`, `sum`, `promote`, `hstack`, `trace`, `linear_op`
63-
- **`elementwise_univariate/`** - Scalar functions applied elementwise: `log`, `exp`, `entr`, `power`, `logistic`, trigonometric, hyperbolic
64-
- **`bivariate/`** - Two-argument operations: `multiply`, `quad_over_lin`, `rel_entr`
65-
- **`other/`** - Special atoms not fitting above categories
61+
- **`affine/`** - Linear operations: `variable`, `constant`, `add`, `neg`, `sum`, `promote`, `hstack`, `trace`, `linear_op`, `index`
62+
- **`elementwise_univariate/`** - Scalar functions applied elementwise: `log`, `exp`, `entr`, `power`, `logistic`, `xexp`, trigonometric (`sin`, `cos`, `tan`), hyperbolic (`sinh`, `tanh`, `asinh`, `atanh`)
63+
- **`bivariate/`** - Two-argument operations: `multiply`, `quad_over_lin`, `rel_entr`, `const_scalar_mult`, `const_vector_mult`, `left_matmul`, `right_matmul`
64+
- **`other/`** - Special atoms: `quad_form`, `prod`
6665

6766
Each atom implements its own `forward`, `jacobian_init`, `eval_jacobian`, and `eval_wsum_hess` functions following a consistent pattern.
6867

@@ -87,7 +86,8 @@ The Python package `dnlp_diff_engine` (in `src/dnlp_diff_engine/`) provides:
8786
**High-level API** (`__init__.py`):
8887
- `C_problem` class wraps the C problem struct
8988
- `convert_problem()` builds expression trees from CVXPY Problem objects
90-
- Atoms are mapped via `ATOM_CONVERTERS` dictionary
89+
- Atoms are mapped via `ATOM_CONVERTERS` dictionary (maps CVXPY atom names → converter functions)
90+
- Special converters handle: matrix multiplication (`_convert_matmul`), multiply with constants (`_convert_multiply`), indexing, reshape (Fortran order only)
9191

9292
**Low-level C extension** (`_core` module, built from `python/bindings.c`):
9393
- Atom constructors: `make_variable`, `make_constant`, `make_log`, `make_exp`, `make_add`, etc.
@@ -108,14 +108,15 @@ Hessian computes weighted sum: `obj_w * H_obj + sum(lambda_i * H_constraint_i)`
108108

109109
## Key Directories
110110

111-
- `include/` - Header files defining public API
112-
- `src/` - C implementation files
113-
- `src/dnlp_diff_engine/` - Python package (installed via pip)
114-
- `python/` - Python bindings C code and binding headers
111+
- `include/` - Header files defining public API (`expr.h`, `problem.h`, atom headers)
112+
- `src/` - C implementation files organized by atom category
113+
- `src/dnlp_diff_engine/` - Python package with high-level API
114+
- `python/` - Python bindings C code (`bindings.c`)
115115
- `python/atoms/` - Python binding headers for each atom type
116116
- `python/problem/` - Python binding headers for problem interface
117+
- `python/tests/` - Python integration tests (run via pytest)
117118
- `tests/` - C tests using minunit framework
118-
- `tests/python/` - Python tests (run via pytest)
119+
- `tests/forward_pass/` - Forward evaluation tests (C)
119120
- `tests/jacobian_tests/` - Jacobian correctness tests (C)
120121
- `tests/wsum_hess/` - Hessian correctness tests (C)
121122

‎include/affine.h‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,4 +18,6 @@ expr *new_trace(expr *child);
1818
expr *new_constant(int d1, int d2, int n_vars, const double *values);
1919
expr *new_variable(int d1, int d2, int var_id, int n_vars);
2020

21+
expr *new_index(expr *child, const int *indices, int n_idxs);
22+
2123
#endif /* AFFINE_H */

‎include/subexpr.h‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,4 +109,13 @@ typedef struct const_vector_mult_expr
109109
double *a; /* length equals node->size */
110110
} const_vector_mult_expr;
111111

112+
/* Index/slicing: y = child[indices] where indices is a list of flat positions */
113+
typedef struct index_expr
114+
{
115+
expr base;
116+
int *indices; /* Flattened indices to select (owned, copied) */
117+
int n_idxs; /* Number of selected elements */
118+
bool has_duplicates; /* True if indices have duplicates (affects Hessian path) */
119+
} index_expr;
120+
112121
#endif /* SUBEXPR_H */

‎python/atoms/index.h‎

Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,53 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
3+
#ifndef ATOM_INDEX_H
4+
#define ATOM_INDEX_H
5+
6+
#include "affine.h"
7+
#include "common.h"
8+
9+
/* Index/slicing: y = child[indices] where indices is a list of flattened positions
10+
*/
11+
static PyObject *py_make_index(PyObject *self, PyObject *args)
12+
{
13+
PyObject *child_capsule;
14+
PyObject *indices_obj;
15+
16+
if (!PyArg_ParseTuple(args, "OO", &child_capsule, &indices_obj))
17+
{
18+
return NULL;
19+
}
20+
21+
expr *child = (expr *) PyCapsule_GetPointer(child_capsule, EXPR_CAPSULE_NAME);
22+
if (!child)
23+
{
24+
PyErr_SetString(PyExc_ValueError, "invalid child capsule");
25+
return NULL;
26+
}
27+
28+
/* Convert indices array to int32 */
29+
PyArrayObject *indices_array = (PyArrayObject *) PyArray_FROM_OTF(
30+
indices_obj, NPY_INT32, NPY_ARRAY_IN_ARRAY);
31+
32+
if (!indices_array)
33+
{
34+
return NULL;
35+
}
36+
37+
int n_idxs = (int) PyArray_SIZE(indices_array);
38+
int *indices_data = (int *) PyArray_DATA(indices_array);
39+
40+
expr *node = new_index(child, indices_data, n_idxs);
41+
42+
Py_DECREF(indices_array);
43+
44+
if (!node)
45+
{
46+
PyErr_SetString(PyExc_RuntimeError, "failed to create index node");
47+
return NULL;
48+
}
49+
expr_retain(node); /* Capsule owns a reference */
50+
return PyCapsule_New(node, EXPR_CAPSULE_NAME, expr_capsule_destructor);
51+
}
52+
53+
#endif /* ATOM_INDEX_H */

‎python/atoms/quad_over_lin.h‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
3+
#ifndef ATOM_QUAD_OVER_LIN_H
4+
#define ATOM_QUAD_OVER_LIN_H
5+
6+
#include "bivariate.h"
7+
#include "common.h"
8+
9+
/* quad_over_lin: y = sum(x^2) / z where x is left, z is right (scalar) */
10+
static PyObject *py_make_quad_over_lin(PyObject *self, PyObject *args)
11+
{
12+
(void)self;
13+
PyObject *left_capsule, *right_capsule;
14+
if (!PyArg_ParseTuple(args, "OO", &left_capsule, &right_capsule))
15+
{
16+
return NULL;
17+
}
18+
expr *left = (expr *)PyCapsule_GetPointer(left_capsule, EXPR_CAPSULE_NAME);
19+
expr *right = (expr *)PyCapsule_GetPointer(right_capsule, EXPR_CAPSULE_NAME);
20+
if (!left || !right)
21+
{
22+
PyErr_SetString(PyExc_ValueError, "invalid child capsule");
23+
return NULL;
24+
}
25+
26+
expr *node = new_quad_over_lin(left, right);
27+
if (!node)
28+
{
29+
PyErr_SetString(PyExc_RuntimeError, "failed to create quad_over_lin node");
30+
return NULL;
31+
}
32+
expr_retain(node); /* Capsule owns a reference */
33+
return PyCapsule_New(node, EXPR_CAPSULE_NAME, expr_capsule_destructor);
34+
}
35+
36+
#endif /* ATOM_QUAD_OVER_LIN_H */

‎python/atoms/rel_entr.h‎

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,36 @@
1+
// SPDX-License-Identifier: Apache-2.0
2+
3+
#ifndef ATOM_REL_ENTR_H
4+
#define ATOM_REL_ENTR_H
5+
6+
#include "bivariate.h"
7+
#include "common.h"
8+
9+
/* rel_entr: rel_entr(x, y) = x * log(x/y) elementwise */
10+
static PyObject *py_make_rel_entr(PyObject *self, PyObject *args)
11+
{
12+
(void)self;
13+
PyObject *left_capsule, *right_capsule;
14+
if (!PyArg_ParseTuple(args, "OO", &left_capsule, &right_capsule))
15+
{
16+
return NULL;
17+
}
18+
expr *left = (expr *)PyCapsule_GetPointer(left_capsule, EXPR_CAPSULE_NAME);
19+
expr *right = (expr *)PyCapsule_GetPointer(right_capsule, EXPR_CAPSULE_NAME);
20+
if (!left || !right)
21+
{
22+
PyErr_SetString(PyExc_ValueError, "invalid child capsule");
23+
return NULL;
24+
}
25+
26+
expr *node = new_rel_entr_vector_args(left, right);
27+
if (!node)
28+
{
29+
PyErr_SetString(PyExc_RuntimeError, "failed to create rel_entr node");
30+
return NULL;
31+
}
32+
expr_retain(node); /* Capsule owns a reference */
33+
return PyCapsule_New(node, EXPR_CAPSULE_NAME, expr_capsule_destructor);
34+
}
35+
36+
#endif /* ATOM_REL_ENTR_H */

‎python/bindings.c‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
#include "atoms/cos.h"
1313
#include "atoms/entr.h"
1414
#include "atoms/exp.h"
15+
#include "atoms/index.h"
1516
#include "atoms/left_matmul.h"
1617
#include "atoms/linear.h"
1718
#include "atoms/log.h"
@@ -21,6 +22,8 @@
2122
#include "atoms/power.h"
2223
#include "atoms/promote.h"
2324
#include "atoms/quad_form.h"
25+
#include "atoms/quad_over_lin.h"
26+
#include "atoms/rel_entr.h"
2427
#include "atoms/right_matmul.h"
2528
#include "atoms/sin.h"
2629
#include "atoms/sinh.h"
@@ -55,6 +58,7 @@ static PyMethodDef DNLPMethods[] = {
5558
{"make_linear", py_make_linear, METH_VARARGS, "Create linear op node"},
5659
{"make_log", py_make_log, METH_VARARGS, "Create log node"},
5760
{"make_exp", py_make_exp, METH_VARARGS, "Create exp node"},
61+
{"make_index", py_make_index, METH_VARARGS, "Create index node"},
5862
{"make_add", py_make_add, METH_VARARGS, "Create add node"},
5963
{"make_sum", py_make_sum, METH_VARARGS, "Create sum node"},
6064
{"make_neg", py_make_neg, METH_VARARGS, "Create neg node"},
@@ -82,6 +86,10 @@ static PyMethodDef DNLPMethods[] = {
8286
"Create right matmul node (f(x) @ A)"},
8387
{"make_quad_form", py_make_quad_form, METH_VARARGS,
8488
"Create quadratic form node (x' * Q * x)"},
89+
{"make_quad_over_lin", py_make_quad_over_lin, METH_VARARGS,
90+
"Create quad_over_lin node (sum(x^2) / y)"},
91+
{"make_rel_entr", py_make_rel_entr, METH_VARARGS,
92+
"Create rel_entr node: x * log(x/y) elementwise"},
8593
{"make_problem", py_make_problem, METH_VARARGS,
8694
"Create problem from objective and constraints"},
8795
{"problem_init_derivatives", py_problem_init_derivatives, METH_VARARGS,

0 commit comments

Comments
 (0)