I have included a reproducible script below showing the numerical impact.
import tvm
from tvm import relay, testing
import numpy as np
prog = \
"""
#[version = "0.0.5"]
def @main(%input_0: Tensor[(5, 10), float32], %output_0: Tensor[(5, 20), float32], %Wrapper_Dense_weight: Tensor[(20, 10), float32], %Wrapper_Dense_bias: Tensor[(20), float32]) -> (float32, Tensor[(5, 20), float32]) {
let %x: (float32, Tensor[(5, 20), float32]) = (
let %x1: Tensor[(5, 20), float32] = nn.dense(%input_0, %Wrapper_Dense_weight, units=20) /* ty=Tensor[(5, 20), float32] */;
let %pred: Tensor[(5, 20), float32] = nn.bias_add(%x1, %Wrapper_Dense_bias, axis=-1) /* ty=Tensor[(5, 20), float32] */;
let %x2: Tensor[(5, 20), float32] = nn.log_softmax(%pred) /* ty=Tensor[(5, 20), float32] */;
let %x3: float32 = nn.cross_entropy_with_logits(%x2, %output_0) /* ty=float32 */;
let %x4: (float32, Tensor[(5, 20), float32]) = (%x3, %pred);
%x4
);
%x
}
"""
inputs = [np.random.randn(*sh).astype('float32') for sh in [(5, 10), (5, 20), (20, 10), (20,)]]
mod = tvm.parser.fromtext(prog)
mod = relay.transform.InferType()(mod)
mod['main'] = relay.transform.gradient(mod['main'], mod=mod, mode='higher_order')
mod = relay.transform.InferType()(mod)
e = relay.create_executor(kind='debug', mod=mod, ctx=tvm.cpu(0), target='llvm').evaluate()
(loss_orig, pred_orig), grad_orig = e(*inputs)
mod = tvm.parser.fromtext(prog)
mod = relay.transform.InferType()(mod)
mod['main'] = relay.transform.gradient(mod['main'], mod=mod, mode='higher_order')
mod = relay.transform.InferType()(mod)
mod = relay.transform.DeadCodeElimination()(mod)
e = relay.create_executor(kind='debug', mod=mod, ctx=tvm.cpu(0), target='llvm').evaluate()
(loss_dce, pred_dce), grad_dce = e(*inputs)
tvm.testing.assert_allclose(loss_dce.asnumpy(), loss_orig.asnumpy())
tvm.testing.assert_allclose(pred_dce.asnumpy(), pred_orig.asnumpy())
for g_dce, g_orig in zip(grad_dce, grad_orig):
tvm.testing.assert_allclose(g_dce.asnumpy(), g_orig.asnumpy())
For now, we should require input modules to DCE not have any references.
The higher-order gradient transformation creates
ref_read's andref_write's which are effectful, and DCE assumes all code is pure. These reference operations will thus get removed, causing gradient updates to be lost. We will likely need to make an alias analysis pass to resolve this. Previously, this may have been unnoticed due to Partial Evaluation being able to often remove all references, but this is not true in general.I have included a reproducible script below showing the numerical impact.
For now, we should require input modules to DCE not have any references.
cc @jroesch @MarisaKirisame