Skip to content

Commit 678716e

Browse files
added capsules in an attempt to prevent segmentation fault
1 parent 5b2ce38 commit 678716e

1 file changed

Lines changed: 15 additions & 14 deletions

File tree

sparse_dot_topn/sparse_dot_topn.pyx

Lines changed: 15 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -172,23 +172,24 @@ cpdef sparse_dot_free(
172172

173173
sparse_dot_free_source(n_row, n_col, Ap, Aj, Ax, Bp, Bj, Bx, lower_bound, Cp, &vCj, &vCx)
174174

175-
cdef np.npy_intp nnz = Cp[n_row]
176-
cdef np.ndarray[np.int32_t, ndim=1] c_indices = np.PyArray_SimpleNewFromData(1, &nnz, np.NPY_INT32, vCj.data())
175+
cdef np.npy_intp nnz[1]
176+
nnz[0] = Cp[n_row]
177+
cdef np.ndarray[np.int32_t, ndim=1] c_indices = np.PyArray_SimpleNewFromData(4, &nnz[0], np.NPY_INT32, vCj.data())
177178
PyArray_ENABLEFLAGS(c_indices, np.NPY_OWNDATA)
178-
cdef np.ndarray[np.double_t, ndim=1] c_data = np.PyArray_SimpleNewFromData(1, &nnz, np.NPY_DOUBLE, vCx.data())
179+
cdef np.ndarray[np.double_t, ndim=1] c_data = np.PyArray_SimpleNewFromData(4, &nnz[0], np.NPY_DOUBLE, vCx.data())
179180
PyArray_ENABLEFLAGS(c_data, np.NPY_OWNDATA)
180181

181-
# cdef const char *name_vCj_capsule = "vCj"
182-
# cdef int* vCj_data = vCj.data()
183-
# vCj_capsule = PyCapsule_New(<void *> vCj_data, name_vCj_capsule, &free_ptr)
184-
# if not PyCapsule_IsValid(vCj_capsule, name_vCj_capsule):
185-
# raise ValueError(f"invalid pointer ({name_vCj_capsule}) to parameters")
186-
#
187-
# cdef const char *name_vCx_capsule = "vCx"
188-
# cdef double* vCx_data = vCx.data()
189-
# vCx_capsule = PyCapsule_New(<void *> vCx_data, name_vCx_capsule, &free_ptr)
190-
# if not PyCapsule_IsValid(vCx_capsule, name_vCx_capsule):
191-
# raise ValueError(f"invalid pointer ({name_vCx_capsule}) to parameters")
182+
cdef const char *name_vCj_capsule = "vCj"
183+
cdef int* vCj_data = vCj.data()
184+
vCj_capsule = PyCapsule_New(<void *> vCj_data, name_vCj_capsule, &free_ptr)
185+
if not PyCapsule_IsValid(vCj_capsule, name_vCj_capsule):
186+
raise ValueError(f"invalid pointer ({name_vCj_capsule}) to parameters")
187+
188+
cdef const char *name_vCx_capsule = "vCx"
189+
cdef double* vCx_data = vCx.data()
190+
vCx_capsule = PyCapsule_New(<void *> vCx_data, name_vCx_capsule, &free_ptr)
191+
if not PyCapsule_IsValid(vCx_capsule, name_vCx_capsule):
192+
raise ValueError(f"invalid pointer ({name_vCx_capsule}) to parameters")
192193

193194
return c_indices, c_data
194195

0 commit comments

Comments
 (0)