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
184 changes: 68 additions & 116 deletions .basedpyright/baseline.json
Original file line number Diff line number Diff line change
Expand Up @@ -2107,14 +2107,6 @@
"lineCount": 1
}
},
{
"code": "reportIncompatibleMethodOverride",
"range": {
"startColumn": 8,
"endColumn": 32,
"lineCount": 1
}
},
{
"code": "reportImplicitOverride",
"range": {
Expand Down Expand Up @@ -8351,14 +8343,6 @@
"lineCount": 1
}
},
{
"code": "reportIncompatibleMethodOverride",
"range": {
"startColumn": 8,
"endColumn": 19,
"lineCount": 1
}
},
{
"code": "reportImplicitOverride",
"range": {
Expand Down Expand Up @@ -13403,6 +13387,62 @@
"lineCount": 1
}
},
{
"code": "reportGeneralTypeIssues",
"range": {
"startColumn": 12,
"endColumn": 24,
"lineCount": 1
}
},
{
"code": "reportUnknownMemberType",
"range": {
"startColumn": 23,
"endColumn": 43,
"lineCount": 1
}
},
{
"code": "reportAttributeAccessIssue",
"range": {
"startColumn": 31,
"endColumn": 43,
"lineCount": 1
}
},
{
"code": "reportUnknownMemberType",
"range": {
"startColumn": 38,
"endColumn": 47,
"lineCount": 1
}
},
{
"code": "reportAttributeAccessIssue",
"range": {
"startColumn": 43,
"endColumn": 47,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 47,
"endColumn": 52,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 64,
"endColumn": 67,
"lineCount": 1
}
},
{
"code": "reportMissingTypeStubs",
"range": {
Expand Down Expand Up @@ -15831,22 +15871,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 18,
"endColumn": 25,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 35,
"endColumn": 42,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
Expand Down Expand Up @@ -16322,10 +16346,10 @@
}
},
{
"code": "reportUnknownArgumentType",
"code": "reportOperatorIssue",
"range": {
"startColumn": 23,
"endColumn": 30,
"startColumn": 18,
"endColumn": 76,
"lineCount": 1
}
},
Expand Down Expand Up @@ -16585,14 +16609,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 38,
"endColumn": 45,
"lineCount": 1
}
},
{
"code": "reportUnknownMemberType",
"range": {
Expand Down Expand Up @@ -16650,15 +16666,15 @@
}
},
{
"code": "reportUnknownMemberType",
"code": "reportOperatorIssue",
"range": {
"startColumn": 48,
"endColumn": 55,
"startColumn": 14,
"endColumn": 72,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"code": "reportUnknownMemberType",
"range": {
"startColumn": 48,
"endColumn": 55,
Expand All @@ -16668,8 +16684,8 @@
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 23,
"endColumn": 30,
"startColumn": 48,
"endColumn": 55,
"lineCount": 1
}
},
Expand Down Expand Up @@ -16826,10 +16842,10 @@
}
},
{
"code": "reportUnknownArgumentType",
"code": "reportOperatorIssue",
"range": {
"startColumn": 23,
"endColumn": 30,
"startColumn": 14,
"endColumn": 72,
"lineCount": 1
}
},
Expand Down Expand Up @@ -16953,14 +16969,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 23,
"endColumn": 30,
"lineCount": 1
}
},
{
"code": "reportMissingParameterType",
"range": {
Expand All @@ -16969,30 +16977,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 23,
"endColumn": 30,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 25,
"endColumn": 32,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 34,
"endColumn": 41,
"lineCount": 1
}
},
{
"code": "reportMissingParameterType",
"range": {
Expand All @@ -17001,14 +16985,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 23,
"endColumn": 30,
"lineCount": 1
}
},
{
"code": "reportUnknownMemberType",
"range": {
Expand Down Expand Up @@ -17073,14 +17049,6 @@
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 23,
"endColumn": 30,
"lineCount": 1
}
},
{
"code": "reportCallIssue",
"range": {
Expand Down Expand Up @@ -17128,22 +17096,6 @@
"endColumn": 47,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 39,
"endColumn": 49,
"lineCount": 1
}
},
{
"code": "reportUnknownArgumentType",
"range": {
"startColumn": 39,
"endColumn": 49,
"lineCount": 1
}
}
],
"./sumpy/test/test_heat_translations.py": [
Expand Down
52 changes: 52 additions & 0 deletions sumpy/symbolic.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,6 +95,9 @@ def _find_symbolic_backend():

# }}}


# {{{ symbolic expressions

if TYPE_CHECKING or not USE_SYMENGINE:
import sympy as sym

Expand Down Expand Up @@ -170,6 +173,8 @@ def doit(expr: Expr) -> Expr:
def unevaluated_pow(a: Expr, b: complex | Expr) -> Expr:
return Pow(a, b, evaluate=False)

# }}}


# {{{ debugging of sympy CSE via Maxima

Expand Down Expand Up @@ -253,6 +258,8 @@ def checked_cse(exprs, symbols=None):
# }}}


# {{{ pymbolic expressions

def sym_real_norm_2(x: Matrix) -> Expr:
return sqrt((x.T*x)[0, 0])

Expand Down Expand Up @@ -295,6 +302,10 @@ def from_sympy(cls, expr: Symbol) -> SpatialConstant:

raise ValueError(f"expression is not a spatial constant: {expr!r}")

# }}}


# {{{ sympy <-> pymbolic interop

class PymbolicToSympyMapper(PymbolicToSympyMapperBase):
def map_spatial_constant(self, expr: SpatialConstant) -> Basic:
Expand Down Expand Up @@ -364,6 +375,45 @@ def map_call(self, expr: prim.Call) -> sym.Basic:
return PymbolicToSympyMapper.map_call(self, expr)


class SympyToPymbolicMapperWithSymbols(SympyToPymbolicMapper):
if USE_SYMENGINE:
@override
def map_Constant(self, expr: object) -> Expression:
Comment thread
inducer marked this conversation as resolved.
if expr is pi:
return prim.Variable("pi")
elif expr is I:
return prim.Variable("I")
else:
return super().map_Constant(expr)
else:
@override
def map_NumberSymbol(self, expr: sym.NumberSymbol) -> Expression:
if expr is pi:
return prim.Variable("pi")
elif expr is I:
return prim.Variable("I")
else:
return super().map_NumberSymbol(expr)

@override
def not_supported(self, expr: object) -> Expression:
if getattr(expr, "is_Function", False):
function_name = self.function_name(expr)
if function_name in {"Hankel1", "BesselJ"}:
order, arg, nderivs = expr.args
if nderivs == 0:
return prim.Variable({
"Hankel1": "hankel_1",
"BesselJ": "bessel_j",
}[function_name])(self.rec(order), self.rec(arg))

return super().not_supported(expr)

# }}}


# {{{ symbolic functions

from sympy import Function as SympyFunction


Expand Down Expand Up @@ -402,4 +452,6 @@ def BesselJ(*args): # ruff:ignore[invalid-function-name]
def Hankel1(*args): # ruff:ignore[invalid-function-name]
return sympify(_SympyHankel1(*args))

# }}}

# vim: fdm=marker
Loading
Loading