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 numba/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,8 +53,8 @@
import numba.core.withcontexts
from numba.core.withcontexts import objmode_context as objmode

# Initialize hardware
import numba.core.extending_hardware
# Initialize target extensions
import numba.core.target_extension

# Initialize typed containers
import numba.typed
Expand Down
4 changes: 2 additions & 2 deletions numba/core/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,9 +234,9 @@ def __init__(self, typing_context, target):

self.address_size = utils.MACHINE_BITS
self.typing_context = typing_context
from numba.core.extending_hardware import hardware_registry
from numba.core.target_extension import target_registry
self.target_name = target
self.target = hardware_registry[target]
self.target = target_registry[target]

# A mapping of installed registries to their loaders
self._registries = {}
Expand Down
10 changes: 1 addition & 9 deletions numba/core/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,6 @@
from numba.core.errors import DeprecationError, NumbaDeprecationWarning
from numba.stencils.stencil import stencil
from numba.core import config, extending, sigutils, registry
from numba.core.extending_hardware import (JitDecorator, hardware_registry,
dispatcher_registry,
resolve_dispatcher_from_str)
from numba.core.registry import TargetRegistry

jit_registry = TargetRegistry()

_logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -191,11 +185,9 @@ def bar(x, y):
return wrapper


# Register the cpu token as using `jit` as the jitter
jit_registry[hardware_registry['cpu']] = jit

def _jit(sigs, locals, target, cache, targetoptions, **dispatcher_args):

from numba.core.target_extension import resolve_dispatcher_from_str
dispatcher = resolve_dispatcher_from_str(target)

def wrapper(func):
Expand Down
6 changes: 3 additions & 3 deletions numba/core/extending.py
Original file line number Diff line number Diff line change
Expand Up @@ -112,16 +112,16 @@ def len_impl(seq):
allow for more control (e.g. disabling non-literal types).

**kwargs prescribes additional arguments passed through to the overload
template. The only accepted key at present is 'hardware' which is a string
corresponding to the hardware that this overload should target.
template. The only accepted key at present is 'target' which is a string
corresponding to the target that this overload should be bound against.
"""
from numba.core.typing.templates import make_overload_template, infer_global

# set default options
opts = _overload_default_jit_options.copy()
opts.update(jit_options) # let user options override

# TODO: abort now if the kwarg 'hardware' relates to unregistered hardware,
# TODO: abort now if the kwarg 'target' relates to an unregistered target,
# this requires sorting out the circular imports first.

def decorate(overload_func):
Expand Down
4 changes: 2 additions & 2 deletions numba/core/ir.py
Original file line number Diff line number Diff line change
Expand Up @@ -418,12 +418,12 @@ def unary(cls, fn, value, loc):
return cls(op=op, loc=loc, fn=fn, value=value)

@classmethod
def call(cls, func, args, kws, loc, vararg=None, hardware=None):
def call(cls, func, args, kws, loc, vararg=None, target=None):
assert isinstance(func, Var)
assert isinstance(loc, Loc)
op = 'call'
return cls(op=op, loc=loc, func=func, args=args, kws=kws,
vararg=vararg, hardware=hardware)
vararg=vararg, target=target)

@classmethod
def build_tuple(cls, items, loc):
Expand Down
5 changes: 2 additions & 3 deletions numba/core/lowering.py
Original file line number Diff line number Diff line change
Expand Up @@ -980,10 +980,9 @@ def _lower_call_normal(self, fnty, expr, signature):
argvals = self.fold_call_args(
fnty, signature, expr.args, expr.vararg, expr.kws,
)
tname = expr.hardware
tname = expr.target
if tname is not None:
from numba.core.extending_hardware import \
resolve_dispatcher_from_str
from numba.core.target_extension import resolve_dispatcher_from_str
disp = resolve_dispatcher_from_str(tname)
hw_ctx = disp.targetdescr.target_context
impl = hw_ctx.get_function(fnty, signature)
Expand Down
16 changes: 8 additions & 8 deletions numba/core/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,30 +73,30 @@ class CPUDispatcher(dispatcher.Dispatcher):
targetdescr = cpu_target


class TargetRegistry(utils.UniqueDict):
class DelayedRegistry(utils.UniqueDict):
"""
A registry of API implementations for various backends.
A unique dictionary but with deferred initialisation of the values.

Attributes
----------
ondemand:

A dictionary of target-name -> function, where function is executed
the first time a target is used. It is used for deferred
initialization for some targets (e.g. gpu).
A dictionary of key -> value, where value is executed
the first time it is is used. It is used for part of a deferred
initialization strategy.
"""
def __init__(self, *args, **kws):
self.ondemand = utils.UniqueDict()
self.key_type = kws.pop('key_type', None)
self.value_type = kws.pop('value_type', None)
self._type_check = self.key_type or self.value_type
super(TargetRegistry, self).__init__(*args, **kws)
super(DelayedRegistry, self).__init__(*args, **kws)

def __getitem__(self, item):
if item in self.ondemand:
self[item] = self.ondemand[item]()
del self.ondemand[item]
return super(TargetRegistry, self).__getitem__(item)
return super(DelayedRegistry, self).__getitem__(item)

def __setitem__(self, key, value):
if self._type_check:
Expand All @@ -109,4 +109,4 @@ def check(x, ty_x):
check(key, self.key_type)
if self.value_type is not None:
check(value, self.value_type)
return super(TargetRegistry, self).__setitem__(key, value)
return super(DelayedRegistry, self).__setitem__(key, value)
70 changes: 41 additions & 29 deletions numba/core/extending_hardware.py → numba/core/target_extension.py
Original file line number Diff line number Diff line change
@@ -1,28 +1,35 @@
from abc import ABC, abstractmethod
from numba.core.registry import TargetRegistry, CPUDispatcher

from numba.core.registry import DelayedRegistry, CPUDispatcher
from numba.core.decorators import jit
from threading import local as tls


_active_context = tls()
_active_context_default = 'cpu'


class HardwareRegistry(TargetRegistry):
class _TargetRegistry(DelayedRegistry):

def __getitem__(self, item):
try:
return super().__getitem__(item)
except KeyError:
msg = "No target is registered against '{}', known targets:\n{}"
known = '\n'.join([f"{k: <{10}} -> {v}"
for k, v in hardware_registry.items()])
for k, v in target_registry.items()])
raise ValueError(msg.format(item, known)) from None


hardware_registry = HardwareRegistry()
# Registry mapping target name strings to Target classes
target_registry = _TargetRegistry()

# Registry mapping Target classes the @jit decorator for that target
jit_registry = DelayedRegistry()


class hardware_target(object):
class target_override(object):
"""Context manager to temporarily override the current target with that
prescribed."""
def __init__(self, name):
self._orig_target = getattr(_active_context, 'target',
_active_context_default)
Expand Down Expand Up @@ -50,9 +57,9 @@ def get_local_target(context):
if len(context.callstack._stack) > 0:
>
else:
class="x x-first x-last">hardware_registry.get(current_target(), None)
class="x x-first x-last">target_registry.get(current_target(), None)
if target is None:
msg = ("The hardware target found is not registered."
msg = ("The target found is not registered."
"Given target was {}.")
raise ValueError(msg.format(target))
else:
Expand All @@ -61,11 +68,11 @@ def get_local_target(context):

def resolve_target_str(target_str):
"""Resolves a target specified as a string to its Target class."""
return hardware_registry[target_str]
return target_registry[target_str]


def resolve_dispatcher_from_str(target_str):
"""Returns the dispatcher associated with a target hardware string"""
"""Returns the dispatcher associated with a target string"""
target_hw = resolve_target_str(target_str)
return dispatcher_registry[target_hw]

Expand All @@ -78,7 +85,7 @@ def __call__(self):


class Target(ABC):
""" Implements a hardware/pseudo-hardware target """
""" Implements a target """

@classmethod
def inherits_from(cls, other):
Expand All @@ -87,46 +94,51 @@ def inherits_from(cls, other):


class Generic(Target):
"""Mark the hardware target as generic, i.e. suitable for compilation on
any target. All hardware must inherit from this.
"""Mark the target as generic, i.e. suitable for compilation on
any target. All must inherit from this.
"""


class CPU(Generic):
"""Mark the hardware target as CPU.
"""Mark the target as CPU.
"""


class GPU(Generic):
"""Mark the hardware target as GPU, i.e. suitable for compilation on a GPU
"""Mark the target as GPU, i.e. suitable for compilation on a GPU
target.
"""


class CUDA(GPU):
"""Mark the hardware target as CUDA.
"""Mark the target as CUDA.
"""


class ROCm(GPU):
"""Mark the hardware target as ROCm.
"""Mark the target as ROCm.
"""


class NPyUfunc(Target):
"""Mark the hardware target as a ufunc
"""Mark the target as a ufunc
"""


hardware_registry['generic'] = Generic
hardware_registry['CPU'] = CPU
hardware_registry['cpu'] = CPU
hardware_registry['GPU'] = GPU
hardware_registry['gpu'] = GPU
hardware_registry['CUDA'] = CUDA
hardware_registry['cuda'] = CUDA
hardware_registry['ROCm'] = ROCm
hardware_registry['npyufunc'] = NPyUfunc
target_registry['generic'] = Generic
target_registry['CPU'] = CPU
target_registry['cpu'] = CPU
target_registry['GPU'] = GPU
target_registry['gpu'] = GPU
target_registry['CUDA'] = CUDA
target_registry['cuda'] = CUDA
target_registry['ROCm'] = ROCm
target_registry['npyufunc'] = NPyUfunc

dispatcher_registry = DelayedRegistry(key_type=Target)


dispatcher_registry = TargetRegistry(key_type=Target)
dispatcher_registry[hardware_registry['cpu']] = CPUDispatcher
# Register the cpu target token with its dispatcher and jit
cpu_target = target_registry['cpu']
dispatcher_registry[cpu_target] = CPUDispatcher
jit_registry[cpu_target] = jit
16 changes: 8 additions & 8 deletions numba/core/types/functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -279,37 +279,37 @@ def get_impl_key(self, sig):
return self._impl_keys[sig.args]

def get_call_type(self, context, args, kws):
from numba.core.extending_hardware import (hardware_registry,
get_local_target)
from numba.core.target_extension import (target_registry,
get_local_target)

prefer_lit = [True, False] # old behavior preferring literal
prefer_not = [False, True] # new behavior preferring non-literal
failures = _ResolutionFailures(context, self, args, kws,
depth=self._depth)

# get the current target hardware
# get the current target target
target_hw = get_local_target(context)

# fish out templates that are specific to the target if a target is
# specified
DEFAULT_HARDWARE = 'generic'
DEFAULT_TARGET = 'generic'
usable = []
for ix, temp_cls in enumerate(self.templates):
# ? Need to do something about this next line
hw = temp_cls.metadata.get('hardware', DEFAULT_HARDWARE)
hw = temp_cls.metadata.get('target', DEFAULT_TARGET)
if hw is not None:
hw_clazz = hardware_registry[hw]
hw_clazz = target_registry[hw]
if target_hw.inherits_from(hw_clazz):
usable.append((temp_cls, hw_clazz, ix))

# sort templates based on hardware specificity
# sort templates based on target specificity
def key(x):
return target_hw.__mro__.index(x[1])
order = [x[0] for x in sorted(usable, key=key)]

if not order:
msg = (f"Function resolution cannot find any matches for function"
f" '{self.key[0]}' for the current hardware: '{target_hw}'.")
f" '{self.key[0]}' for the current target: '{target_hw}'.")
raise errors.UnsupportedError(msg)

self._depth += 1
Expand Down
Loading