Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions arrayfire_wrapper/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
__arrayfire_version__ = ARRAYFIRE_VERSION

__all__ += ["add"]
from arrayfire_wrapper.library.mathematical_functions import add
from arrayfire_wrapper.lib.mathematical_functions import add

__all__ += ["randu"]
from arrayfire_wrapper.library.create_and_modify_array import randu
from arrayfire_wrapper.lib.create_and_modify_array import randu
4 changes: 0 additions & 4 deletions arrayfire_wrapper/_typing.py

This file was deleted.

9 changes: 2 additions & 7 deletions arrayfire_wrapper/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,14 +10,9 @@
from pathlib import Path
from typing import Iterator

from arrayfire_wrapper.version import ARRAYFIRE_VER_MAJOR

from ._logger import logger


def is_arch_x86() -> bool:
machine = platform.machine()
return platform.architecture()[0][0:2] == "32" and (machine[-2:] == "86" or machine[0:3] == "arm")
from .defines import is_arch_x86
from .version import ARRAYFIRE_VER_MAJOR


class _LibPrefixes(Enum):
Expand Down
58 changes: 58 additions & 0 deletions arrayfire_wrapper/defines.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
from __future__ import annotations

import ctypes
import platform
from dataclasses import dataclass
from enum import Enum
from typing import Type


def is_arch_x86() -> bool:
machine = platform.machine()
return platform.architecture()[0][0:2] == "32" and (machine[-2:] == "86" or machine[0:3] == "arm")


# A handle for an internal array object
class AFArray(ctypes.c_void_p):
@classmethod
def create_null_pointer(cls) -> AFArray:
cls.value = None
return cls()


CType = Type[ctypes._SimpleCData]
CDimT = ctypes.c_int if is_arch_x86() else ctypes.c_longlong


@dataclass(frozen=True)
class ArrayBuffer:
address: int
length: int = 0


class CShape(tuple):
def __new__(cls, *args: int) -> CShape:
cls.original_shape = len(args)
return tuple.__new__(cls, args)

def __init__(self, x1: int = 1, x2: int = 1, x3: int = 1, x4: int = 1) -> None:
self.x1 = x1
self.x2 = x2
self.x3 = x3
self.x4 = x4

def __repr__(self) -> str:
return f"{self.__class__.__name__}{self.x1, self.x2, self.x3, self.x4}"

@property
def c_array(self): # type: ignore[no-untyped-def]
c_shape = CDimT * 4 # ctypes.c_int | ctypes.c_longlong * 4
return c_shape(CDimT(self.x1), CDimT(self.x2), CDimT(self.x3), CDimT(self.x4))


class Moment(Enum):
M00 = 1
M01 = 2
M10 = 4
M11 = 8
FIRST_ORDER = M00 | M01 | M10 | M11
30 changes: 2 additions & 28 deletions arrayfire_wrapper/dtypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,9 @@

import ctypes
from dataclasses import dataclass
from typing import Type, TypeAlias

from .backend import is_arch_x86
from .defines import CType

CType = Type[ctypes._SimpleCData]
_python_bool = bool


Expand Down Expand Up @@ -63,31 +61,7 @@ def is_complex_dtype(dtype: Dtype) -> _python_bool:
return dtype in {complex64, complex128}


c_dim_t = ctypes.c_int if is_arch_x86() else ctypes.c_longlong
ShapeType = tuple[int, ...]


class CShape(tuple):
def __new__(cls, *args: int) -> CShape:
cls.original_shape = len(args)
return tuple.__new__(cls, args)

def __init__(self, x1: int = 1, x2: int = 1, x3: int = 1, x4: int = 1) -> None:
self.x1 = x1
self.x2 = x2
self.x3 = x3
self.x4 = x4

def __repr__(self) -> str:
return f"{self.__class__.__name__}{self.x1, self.x2, self.x3, self.x4}"

@property
def c_array(self): # type: ignore[no-untyped-def]
c_shape = c_dim_t * 4 # ctypes.c_int | ctypes.c_longlong * 4
return c_shape(c_dim_t(self.x1), c_dim_t(self.x2), c_dim_t(self.x3), c_dim_t(self.x4))


def to_str(c_str: ctypes.c_char_p) -> str:
def to_str(c_str: ctypes.c_char_p | ctypes.Array[ctypes.c_char]) -> str:
return str(c_str.value.decode("utf-8")) # type: ignore[union-attr]


Expand Down
15 changes: 15 additions & 0 deletions arrayfire_wrapper/lib/_broadcast.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
class Bcast:
def __init__(self) -> None:
self._flag: bool = False

def get(self) -> bool:
return self._flag

def set(self, flag: bool) -> None:
self._flag = flag

def toggle(self) -> None:
self._flag ^= True


bcast_var: Bcast = Bcast()
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,8 @@
from enum import Enum

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.dtypes import c_dim_t, to_str
from arrayfire_wrapper.defines import CDimT
from arrayfire_wrapper.dtypes import to_str


class _ErrorCodes(Enum):
Expand All @@ -14,6 +15,6 @@ def safe_call(c_err: int) -> None:
return

err_str = ctypes.c_char_p(0)
err_len = c_dim_t(0)
err_len = CDimT(0)
_backend.clib.af_get_last_error(ctypes.pointer(err_str), ctypes.pointer(err_len))
raise RuntimeError(to_str(err_str))
Original file line number Diff line number Diff line change
Expand Up @@ -3,9 +3,9 @@
import ctypes
from collections.abc import Callable

from arrayfire_wrapper._typing import AFArray
from arrayfire_wrapper.library._broadcast import bcast_var
from arrayfire_wrapper.library._error_handler import safe_call
from arrayfire_wrapper.defines import AFArray
from arrayfire_wrapper.lib._broadcast import bcast_var
from arrayfire_wrapper.lib._error_handler import safe_call


def binary_op(c_func: Callable, lhs: AFArray, rhs: AFArray, /) -> AFArray:
Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
# flake8: noqa
__all__ = [
"AFRandomEngine",
"AFRandomEngineHandle",
"create_random_engine",
"random_engine_get_seed",
"random_engine_get_type",
Expand All @@ -12,7 +12,7 @@
]

from .create_array.random_number_generation import (
AFRandomEngine,
AFRandomEngineHandle,
create_random_engine,
random_engine_get_seed,
random_engine_get_type,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
# flake8: noqa
__all__ = ["assign_gen", "assign_seq"]

from .assign import assign_gen, assign_seq

__all__ += ["lookup"]

from .lookup import lookup
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
import ctypes
from typing import Any

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.defines import AFArray

from ..._error_handler import safe_call

# TODO fix typing for indices across all functions


def assign_gen(lhs: AFArray, rhs: AFArray, ndims: int, indices: Any, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__index__func__assign.htm#ga93cd5199c647dce0e3b823f063b352ae
"""
out = AFArray(0)
safe_call(_backend.clib.af_assign_gen(ctypes.pointer(out), lhs, ndims, indices.pointer, rhs))
return out


def assign_seq(lhs: AFArray, rhs: AFArray, ndims: int, indices: Any, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__index__func__assign.htm#ga3b201c3114941b6f8d0e344afcd18457
"""
out = AFArray(0)
safe_call(_backend.clib.af_assign_seq(ctypes.pointer(out), lhs, ndims, indices.pointer, rhs))
return out
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,9 @@
from typing import Any

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.library._broadcast import bcast_var
from arrayfire_wrapper.lib._broadcast import bcast_var

from ._error_handler import safe_call
from ..._error_handler import safe_call


class _IndexSequence(ctypes.Structure):
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
import ctypes

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.defines import AFArray

from ..._error_handler import safe_call


def lookup(arr: AFArray, indices: AFArray, dim: int, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__index__func__assign.htm#ga93cd5199c647dce0e3b823f063b352ae
"""
out = AFArray(0)
safe_call(_backend.clib.af_assign_gen(ctypes.pointer(out), arr, indices, ctypes.c_int(dim)))
return out
Original file line number Diff line number Diff line change
@@ -0,0 +1,54 @@
# flake8: noqa
__all__ = ["constant", "constant_complex", "constant_long", "constant_ulong"]

from .constant import constant, constant_complex, constant_long, constant_ulong

__all__ += [
"create_random_engine",
"random_engine_get_seed",
"random_engine_get_type",
"random_engine_set_seed",
"random_engine_set_type",
"random_uniform",
"randu",
"release_random_engine",
]

from .random_number_generation import (
create_random_engine,
random_engine_get_seed,
random_engine_get_type,
random_engine_set_seed,
random_engine_set_type,
random_uniform,
randu,
release_random_engine,
)

__all__ += ["diag_create", "diag_extract"]

from .diag import diag_create, diag_extract

__all__ += ["identity"]

from .identity import identity

__all__ += ["iota"]

from .iota import iota

__all__ += ["lower"]

from .lower import lower

__all__ += ["pad"]

from .pad import pad

__all__ += ["range"]

from .range import range

__all__ += ["upper"]

from .upper import upper
Original file line number Diff line number Diff line change
@@ -1,20 +1,17 @@
import ctypes
from typing import TYPE_CHECKING

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.dtypes import CShape, Dtype
from arrayfire_wrapper.defines import AFArray, CShape
from arrayfire_wrapper.dtypes import Dtype

from ..._error_handler import safe_call

if TYPE_CHECKING:
from arrayfire_wrapper._typing import AFArray


def constant(number: int | float, shape: tuple[int, ...], dtype: Dtype, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__data__func__constant.htm#gafc51b6a98765dd24cd4139f3bde00670
"""
out = ctypes.c_void_p(0)
out = AFArray.create_null_pointer()
c_shape = CShape(*shape)

safe_call(
Expand All @@ -29,7 +26,7 @@ def constant_complex(number: int | float | complex, shape: tuple[int, ...], dtyp
"""
source: https://arrayfire.org/docs/group__data__func__constant.htm#ga5a083b1f3cd8a72a41f151de3bdea1a2
"""
out = ctypes.c_void_p(0)
out = AFArray.create_null_pointer()
c_shape = CShape(*shape)

safe_call(
Expand All @@ -49,7 +46,7 @@ def constant_long(number: int | float, shape: tuple[int, ...], dtype: Dtype, /)
"""
source: https://arrayfire.org/docs/group__data__func__constant.htm#ga10f1c9fad1ce9e9fefd885d5a1d1fd49
"""
out = ctypes.c_void_p(0)
out = AFArray.create_null_pointer()
c_shape = CShape(*shape)

safe_call(
Expand All @@ -64,7 +61,9 @@ def constant_ulong(number: int | float, shape: tuple[int, ...], dtype: Dtype, /)
"""
source: https://arrayfire.org/docs/group__data__func__constant.htm#ga67af670cc9314589f8134019f5e68809
"""
out = ctypes.c_void_p(0)
# out = ctypes.c_void_p(0)
# out = AFArray(0)
out = AFArray.create_null_pointer()
c_shape = CShape(*shape)

safe_call(
Expand Down
24 changes: 24 additions & 0 deletions arrayfire_wrapper/lib/create_and_modify_array/create_array/diag.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,24 @@
import ctypes

from arrayfire_wrapper.backend import _backend
from arrayfire_wrapper.defines import AFArray

from ..._error_handler import safe_call


def diag_create(arr: AFArray, num: int, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__data__func__diag.htm#gaecc9950acc89aefcb99ad805af8aa29b
"""
out = AFArray(0)
safe_call(_backend.clib.af_diag_create(ctypes.pointer(out), arr, ctypes.c_int(num)))
return out


def diag_extract(arr: AFArray, num: int, /) -> AFArray:
"""
source: https://arrayfire.org/docs/group__data__func__diag.htm#ga0a28a19534f3c92f11373d662c183061
"""
out = AFArray(0)
safe_call(_backend.clib.af_diag_extract(ctypes.pointer(out), arr, ctypes.c_int(num)))
return out
Loading