Repository navigation
[Relay][RFC] Compilation for heterogeneous execution #2296
Description
Activity
Supporting heterogeneous execution for Relay is great!
I have a couple ideas on the implementation though.
In general I think we should avoid modifying Relay's AST if possible.
I'm motivated by a couple simple design choices. If we extend the AST for every feature the AST will quickly become large and complex. The AST is part of the user facing interface (where users are compiler/tool authors) and modifications are transparent to them, requiring us to be careful about backwards compatibility.
One design the Relay team has stumbled upon is using synthetic operators (i.e
on_device(expr, device_id=n)).In this design the user and/or frontend may call the
on_deviceoperator, its first argument is an arbitrary expression, the result of calling this operator schedules the computation on the device with the idn.The annotation pass can simply rewrite the program to place the correct calls to
on_device, and then during compilation we can rewrite them into the correct communication operations.We should also make the modifications to the planning API as suggested.
One good property about this approach is we can easily try out different heterogenous execution strategies in libraries, as well as enabling other backends (such as the 2.0 runtime we have an open RFC on) to compile the
on_deviceoperator appropriately.@jroesch Thanks for your suggestion. I am not sure if I fully understand what you said. Let's me walk through a simple example.
x = relay.var("x") y = relay.var("y") z = relay.var("z") add = relay.add(x, y) sub = relay.sub(add, z) ...Now we what to assign
addandsubto device0and1, respectively. Do you mean that users want toon_devicelike the following?x = relay.var("x") y = relay.var("y") z = relay.var("z") add = relay.add(x, y) on_device(add, 0) sub = relay.sub(add, z) on_device(sub, 1) ...Then we run the annotation pass to rewrite the program to:
x = relay.var("x") y = relay.var("y") z = relay.var("z") add = relay.add(x, y) copy0 = relay.device_copy(add) sub = relay.sub(copy0, z) copy1 = relay.device_copy(sub) ...Please correct me if there is any misunderstanding. Otherwise, I have a couple of questions:
- Are users supposed to go through the network and perform the annotation? It looks this might be handy when the network is large.
- How should we propagate the device ids? I think one way is to propagate the device id information based in a bottom-up manner where the device_copy op should have an
device_idattribute.
@zhiics can you update the proposal to reflect the latest result of discussion? I see there is an on_device annotation to mark the region of computing.
Let us also designate a specific namespace for such annotation, we will likely have a few of them. Example candidates are:
- relay.opt.on_device
- relay.annotate.on_device
RFC for namespace convention #2391
closed by #2361
Motivation
The graph runtime is now able to support heterogeneous execution through annotation with various device ids. It is important to have a compiler pass to enable annotation from the frontend so that users have the flexibility to annotate the operators with "the best" device. This RFC proposes to add annotation in Relay as a standalone pass. There prototype implementation is here.
Action items
Some design items are listed as following:
fallbackattribute to indicate if it will fallback.CallNodeis attached with adevice_idattribute to indicate which device it should be annotated to (by default it is 0, meaning no annotation is required)Dict[op_name, device]map tobuild, or enable fallback by addingset_fallbackto an operator. More sophisticated annotation schemes (i.e. the ones with cost functions by taking device communication and data transferring overhead into account) could be explored in the future.fcomputeandfschedule. These ops could be omitted during lowering as well since the real data copy will be performed during runtime.Proposed APIs
buildAPI is like the following:def build(func, target=None, target_host=None, params=None, op_name_device=None, fallback_device=None):.where heterogeneous compilation is enabled when target is a dict of device to target.
def annotate_ops(expr, op_name_dev_map, fallback_device):During annotation, the
device_idof aCallNodeis set tofallback_deviceif its operator is registered withfallbackor it is not explicitly specified where it should be allocated to in the map.PlanAPI ingraph_plan_memoryneeds to be changed slightly. Now in addition to returning a list ofstorage_id, the correspondingdevice_idalso has to be returned.