Skip to content

Commit 3bfec33

Browse files
Fix: overloads for dolfinx.fem.form (#4419)
* Fix: overloads for dolfinx.fem.form * Apply suggestions from code review Co-authored-by: Paul T. Kühner <56360279+schnellerhase@users.noreply.github.com> * Implicaitons * more * another * .. * snowball * . * ig * ig * Align constants with coeffs
1 parent 0d03e5c commit 3bfec33

3 files changed

Lines changed: 47 additions & 37 deletions

File tree

python/dolfinx/fem/assemble.py

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -25,16 +25,20 @@
2525
from dolfinx.fem.utils import create_sparsity_pattern
2626

2727

28+
@typing.overload
29+
def pack_constants(form: None) -> None: ...
30+
31+
2832
@typing.overload
2933
def pack_constants(form: Form) -> npt.NDArray: ...
3034

3135

3236
@typing.overload
33-
def pack_constants(form: Sequence[Form]) -> list[npt.NDArray]: ...
37+
def pack_constants(form: Sequence[Form | None]) -> list[npt.NDArray]: ...
3438

3539

3640
def pack_constants(
37-
form: Form | Sequence[Form] | None,
41+
form: Form | Sequence[Form | None] | None,
3842
) -> npt.NDArray | list[npt.NDArray] | None:
3943
"""Pack form constants for use in assembly.
4044
@@ -56,7 +60,7 @@ def pack_constants(
5660
if form is None:
5761
return None
5862
elif isinstance(form, Sequence):
59-
return list(map(pack_constants, form))
63+
return list(map(pack_constants, form)) # type: ignore
6064
else:
6165
return _pack_constants(form._cpp_object)
6266

@@ -67,12 +71,12 @@ def pack_coefficients(form: Form | None) -> dict[tuple[IntegralType, int], npt.N
6771

6872
@typing.overload
6973
def pack_coefficients(
70-
form: Sequence[Form],
74+
form: Sequence[Form | None],
7175
) -> list[dict[tuple[IntegralType, int], npt.NDArray]]: ...
7276

7377

7478
def pack_coefficients(
75-
form: Form | Sequence[Form] | None,
79+
form: Form | Sequence[Form | None] | None,
7680
) -> (
7781
dict[tuple[IntegralType, int], npt.NDArray] | list[dict[tuple[IntegralType, int], npt.NDArray]]
7882
):

python/dolfinx/fem/forms.py

Lines changed: 17 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -343,48 +343,48 @@ def mixed_topology_form(
343343

344344
@typing.overload
345345
def form(
346-
form: Sequence[Sequence[ufl.Form]],
346+
form: None,
347347
dtype: npt.DTypeLike = default_scalar_type,
348348
form_compiler_options: dict | None = None,
349349
jit_options: dict | None = None,
350350
jit_comm: MPI.Intracomm | None = None,
351351
entity_maps: Sequence[_EntityMap] | None = None,
352-
) -> list[list[Form]]: ...
352+
) -> None: ...
353353
@typing.overload
354354
def form(
355-
form: Sequence[ufl.Form],
355+
form: ufl.Form,
356356
dtype: npt.DTypeLike = default_scalar_type,
357357
form_compiler_options: dict | None = None,
358358
jit_options: dict | None = None,
359359
jit_comm: MPI.Intracomm | None = None,
360360
entity_maps: Sequence[_EntityMap] | None = None,
361-
) -> list[Form]: ...
361+
) -> Form: ...
362362
@typing.overload
363363
def form(
364-
form: None,
364+
form: Sequence[ufl.Form | None],
365365
dtype: npt.DTypeLike = default_scalar_type,
366366
form_compiler_options: dict | None = None,
367367
jit_options: dict | None = None,
368368
jit_comm: MPI.Intracomm | None = None,
369369
entity_maps: Sequence[_EntityMap] | None = None,
370-
) -> None: ...
370+
) -> list[Form | None]: ...
371371
@typing.overload
372372
def form(
373-
form: ufl.Form,
373+
form: Sequence[Sequence[ufl.Form | None]],
374374
dtype: npt.DTypeLike = default_scalar_type,
375375
form_compiler_options: dict | None = None,
376376
jit_options: dict | None = None,
377377
jit_comm: MPI.Intracomm | None = None,
378378
entity_maps: Sequence[_EntityMap] | None = None,
379-
) -> Form: ...
379+
) -> list[list[Form | None]]: ...
380380
def form(
381-
form: ufl.Form | Sequence[ufl.Form] | Sequence[Sequence[ufl.Form]] | None,
381+
form: ufl.Form | Sequence[ufl.Form | None] | Sequence[Sequence[ufl.Form | None]] | None,
382382
dtype: npt.DTypeLike = default_scalar_type,
383383
form_compiler_options: dict | None = None,
384384
jit_options: dict | None = None,
385385
jit_comm: MPI.Intracomm | None = None,
386386
entity_maps: Sequence[_EntityMap] | None = None,
387-
) -> Form | list[Form] | list[list[Form]] | None:
387+
) -> Form | list[Form | None] | list[list[Form | None]] | None:
388388
"""Create a Form or list of Forms.
389389
390390
Args:
@@ -516,7 +516,7 @@ def _zero_form(form: ufl.ZeroBaseForm) -> Form:
516516
return Form(f)
517517

518518
def _create_form(
519-
form: ufl.Form | Sequence[ufl.Form] | Sequence[Sequence[ufl.Form]] | None,
519+
form: ufl.Form | Sequence[ufl.Form | None] | Sequence[Sequence[ufl.Form | None]] | None,
520520
) -> typing.Any:
521521
"""Recursively convert ufl.Forms to dolfinx.fem.Form.
522522
@@ -536,21 +536,23 @@ def _create_form(
536536
else:
537537
return form
538538

539-
return typing.cast("Form | list[Form] | list[list[Form]] | None", _create_form(form))
539+
return typing.cast(
540+
"Form | list[Form | None] | list[list[Form | None]] | None", _create_form(form)
541+
)
540542

541543

542544
@typing.overload
543545
def extract_function_spaces(forms: Form, index: None = None) -> FunctionSpace | None: ...
544546
@typing.overload
545547
def extract_function_spaces(
546-
forms: Sequence[Form], index: None = None
548+
forms: Sequence[Form | None], index: None = None
547549
) -> list[FunctionSpace | None]: ...
548550
@typing.overload
549551
def extract_function_spaces(
550-
forms: Sequence[Sequence[Form]], index: int = 0
552+
forms: Sequence[Sequence[Form | None]], index: int = 0
551553
) -> list[FunctionSpace | None]: ...
552554
def extract_function_spaces(
553-
forms: Form | Sequence[Form] | Sequence[Sequence[Form]], index: int | None = None
555+
forms: Form | Sequence[Form | None] | Sequence[Sequence[Form | None]], index: int | None = None
554556
) -> FunctionSpace | list[FunctionSpace | None] | None:
555557
"""Extract common function spaces from an array of forms.
556558

python/dolfinx/fem/petsc.py

Lines changed: 21 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -153,7 +153,7 @@ def create_vector(
153153

154154

155155
def create_matrix(
156-
a: Form | Sequence[Sequence[Form]],
156+
a: Form | Sequence[Sequence[Form | None]],
157157
kind: str | Sequence[Sequence[str]] | None = None,
158158
) -> PETSc.Mat:
159159
"""Create a matrix compatible with a sequence of bilinear forms.
@@ -381,7 +381,7 @@ def _assemble_vector_petsc(
381381
# -- Matrix assembly ------------------------------------------------------
382382
@overload
383383
def assemble_matrix(
384-
a: Form | Sequence[Sequence[Form]],
384+
a: Form | Sequence[Sequence[Form | None]],
385385
bcs: Sequence[DirichletBC] | None = None,
386386
diag: float = 1.0,
387387
constants: npt.NDArray | Sequence[Sequence[npt.NDArray]] | None = None,
@@ -395,7 +395,7 @@ def assemble_matrix(
395395
@overload
396396
def assemble_matrix(
397397
A: PETSc.Mat,
398-
a: Form | Sequence[Sequence[Form]],
398+
a: Form | Sequence[Sequence[Form | None]],
399399
bcs: Sequence[DirichletBC] | None = None,
400400
diag: float = 1.0,
401401
constants: npt.NDArray | Sequence[Sequence[npt.NDArray]] | None = None,
@@ -409,7 +409,7 @@ def assemble_matrix(
409409

410410
@functools.singledispatch
411411
def assemble_matrix(
412-
a: Form | Sequence[Sequence[Form]],
412+
a: Form | Sequence[Sequence[Form | None]],
413413
bcs: Sequence[DirichletBC] | None = None,
414414
diag: float = 1,
415415
constants: npt.NDArray | Sequence[Sequence[npt.NDArray]] | None = None,
@@ -478,7 +478,7 @@ def assemble_matrix(
478478
@assemble_matrix.register # type: ignore[attr-defined]
479479
def _assemble_matrix_petsc(
480480
A: PETSc.Mat,
481-
a: Form | Sequence[Sequence[Form]],
481+
a: Form | Sequence[Sequence[Form | None]],
482482
bcs: Sequence[DirichletBC] | None = None,
483483
diag: float = 1,
484484
constants: npt.NDArray | Sequence[Sequence[npt.NDArray]] | None = None,
@@ -513,12 +513,12 @@ def _assemble_matrix_petsc(
513513
if a_block is not None:
514514
Asub = A.getNestSubMatrix(i, j)
515515
_assemble_matrix_petsc(Asub, a_block, bcs, diag, const, coeff)
516-
elif i == j:
516+
elif i == j and bcs is not None:
517517
for bc in bcs:
518518
row_forms = [row_form for row_form in a_row if row_form is not None]
519519
if len(row_forms) == 0:
520520
raise ValueError(f"Row {i} of forms is entirely 'None'.")
521-
if row_forms[0].function_spaces[0].contains(bc.function_space._cpp_object):
521+
if row_forms[0].function_spaces[0].contains(bc.function_space._cpp_object): # type: ignore
522522
raise RuntimeError(
523523
f"Diagonal sub-block ({i}, {j}) cannot be 'None'"
524524
" and have DirichletBC applied."
@@ -559,7 +559,7 @@ def _assemble_matrix_petsc(
559559
)
560560
A.restoreLocalSubMatrix(is0[i], is1[j], Asub)
561561
elif i == j:
562-
for bc in _bcs:
562+
for bc in _bcs: # type: ignore
563563
row_forms = [row_form for row_form in a_row if row_form is not None]
564564
if len(row_forms) == 0:
565565
raise ValueError(f"Row {i} of forms is entirely 'None'.")
@@ -599,7 +599,7 @@ def _assemble_matrix_petsc(
599599

600600
def apply_lifting(
601601
b: PETSc.Vec,
602-
a: Sequence[Form] | Sequence[Sequence[Form]],
602+
a: Sequence[Form | None] | Sequence[Sequence[Form | None]],
603603
bcs: Sequence[DirichletBC] | Sequence[Sequence[DirichletBC]] | None,
604604
x0: Sequence[PETSc.Vec] | None = None,
605605
alpha: float = 1,
@@ -665,12 +665,16 @@ def apply_lifting(
665665
"""
666666
if b.getType() == PETSc.Vec.Type.NEST:
667667
x0 = [] if x0 is None else x0.getNestSubVecs() # type: ignore[attr-defined]
668-
constants = [pack_constants(forms) for forms in a] if constants is None else constants # type: ignore[assignment]
669-
coeffs = [pack_coefficients(forms) for forms in a] if coeffs is None else coeffs # type: ignore[misc]
668+
if constants is None:
669+
constants = [pack_constants(forms) for forms in a] # type: ignore
670+
if coeffs is None:
671+
coeffs = [pack_coefficients(forms) for forms in a] # type: ignore
672+
assert coeffs is not None
673+
assert constants is not None
670674
for b_sub, a_sub, const, coeff in zip(
671675
b.getNestSubVecs(),
672676
a,
673-
constants, # type: ignore[arg-type]
677+
constants,
674678
coeffs,
675679
strict=True,
676680
):
@@ -697,8 +701,8 @@ def apply_lifting(
697701
for i, (a_, off0, off1, offg0, offg1) in enumerate(
698702
zip(a, offset0[:-1], offset0[1:], offset1[:-1], offset1[1:], strict=True)
699703
):
700-
const = pack_constants(a_) if constants is None else constants[i] # type: ignore[call-overload]
701-
coeff = pack_coefficients(a_) if coeffs is None else coeffs[i] # type: ignore[index, call-overload, assignment]
704+
const = pack_constants(a_) if constants is None else constants[i] # type: ignore
705+
coeff = pack_coefficients(a_) if coeffs is None else coeffs[i] # type: ignore
702706
const_ = [
703707
np.empty(0, dtype=PETSc.ScalarType) if val is None else val
704708
for val in const
@@ -1085,7 +1089,7 @@ def a(self) -> Form | Sequence[Sequence[Form]]:
10851089
return typing.cast(Form | Sequence[Sequence[Form]], self._a)
10861090

10871091
@property
1088-
def preconditioner(self) -> Form | Sequence[Sequence[Form]] | None:
1092+
def preconditioner(self) -> Form | Sequence[Sequence[Form | None]] | None:
10891093
"""The compiled bilinear form representing the preconditioner."""
10901094
return self._preconditioner
10911095

@@ -1280,7 +1284,7 @@ class NonlinearProblem(typing.Generic[_U]):
12801284
""" # noqa: D301
12811285

12821286
_P_mat: PETSc.Mat | None
1283-
_preconditioner: Form | Sequence[Sequence[Form]] | None
1287+
_preconditioner: Form | Sequence[Sequence[Form | None]] | None
12841288

12851289
@typing.overload
12861290
def __init__(
@@ -1539,7 +1543,7 @@ def J(self) -> Form | Sequence[Sequence[Form]]:
15391543
return typing.cast(Form | Sequence[Sequence[Form]], self._J)
15401544

15411545
@property
1542-
def preconditioner(self) -> Form | Sequence[Sequence[Form]] | None:
1546+
def preconditioner(self) -> Form | Sequence[Sequence[Form | None]] | None:
15431547
"""The compiled preconditioner."""
15441548
return self._preconditioner
15451549

0 commit comments

Comments
 (0)