sciml.interpreter¶
The interpreter of the forward pass of a network.
The forward pass of the NN YAML is the list of the nodes of a torch.fx
graph: placeholder (an input), call_module (a layer), call_function and
call_method (a function), and output. evaluate walks the list once with
a backend, so the forward pass in numpy and the compilation into expressions
are the same interpreter and the same layers.
The NN YAML writes a reference to a node as the name of the node, so a
reference and a string literal look the same. The arguments of a layer and of
an output are references only. The arguments of a function are references
where the implementation takes arrays, i.e. where its parameter is annotated
with numpy.ndarray (or a type which holds it, e.g. a sequence of arrays),
and literals everywhere else, so gelu(x, "tanh") stays the literal "tanh"
in a network with a node named tanh.
evaluate
¶
Evaluate the forward pass of a network.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
model
|
NNModel
|
the architecture of the network. |
required |
parameters
|
Mapping[str, Mapping[str, ndarray]]
|
the arrays of the layers, layer id -> array name -> array, in the PyTorch layout. |
required |
inputs
|
Sequence[ArrayLike]
|
the inputs, one per |
required |
backend
|
Backend
|
the backend the layers and functions are evaluated with. |
required |
on_node
|
NodeHook | None
|
called after every layer and function with the node and its value, the next nodes see what it returns. The compilation uses it to replace the expressions of a node by symbols. |
None
|
Returns:
| Type | Description |
|---|---|
ndarray
|
The outputs of the network, arrays which share no memory with the |
...
|
inputs. |
Raises:
| Type | Description |
|---|---|
UnsupportedLayerError
|
if a layer or function is not implemented, not available in the backend or called with arguments its implementation does not take. |
ValueError
|
if the number of inputs is not the number of placeholders, if a node has an unknown opcode, if an argument does not name a node which was evaluated or if a layer or function cannot be evaluated on its input; the message names the network and the node. |