Skip to content

[Relay] Dead Code Elimination is unsound with references #6803

Description

@altanh

The higher-order gradient transformation creates ref_read's and ref_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.

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.

cc @jroesch @MarisaKirisame

Metadata

Metadata

Assignees

No one assigned

    Labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions