Skip to content
Open
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
60 changes: 54 additions & 6 deletions backends/transforms/addmm_mm_to_linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,14 @@
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from typing import Optional

import torch
from executorch.exir.dialects._ops import ops as exir_ops
from executorch.exir.pass_base import ExportPass, PassResult

from executorch.exir.sym_util import eval_shape, eval_shape_upper_bound

from torch._export.utils import is_buffer, is_lifted_tensor_constant, is_param
from torch.export import ExportedProgram

_int64_max_dim_val = torch.iinfo(torch.int64).max - 1

Expand All @@ -32,6 +34,33 @@ def get_shape(input_node: torch.fx.Node):
return upper_bound_shape


def is_constant_tensor(
node: torch.fx.Node, exported_program: Optional[ExportedProgram]
) -> bool:
"""
Whether `node` produces a tensor whose contents are known at build time.

`mm`/`addmm` place no constraint on their second operand, but a `linear`
node's weight does carry one: backends prepack it while building the
delegate graph. Rewriting to `linear` is therefore only valid when the
operand really is a constant.
"""
if node.op == "get_attr":
return True
if node.op != "placeholder":
return False
if exported_program is None:
# Without the owning program a lifted parameter cannot be told apart
# from a user input. Placeholders were always rewritten before, so keep
# accepting them rather than regressing callers that pass no program.
return True
return (
is_param(exported_program, node)
or is_buffer(exported_program, node)
or is_lifted_tensor_constant(exported_program, node)
)


def get_dqlinear_input(node: torch.fx.Node):
ops = exir_ops.edge
node_to_backtrack = node
Expand Down Expand Up @@ -99,7 +128,9 @@ def replace_linear_view_copy_input_output(graph: torch.fx.Graph) -> torch.fx.Gra
return graph


def replace_addmm_mm_with_linear(graph: torch.fx.Graph) -> torch.fx.Graph:
def replace_addmm_mm_with_linear(
graph: torch.fx.Graph, exported_program: Optional[ExportedProgram] = None
) -> torch.fx.Graph:
"""
Replace calls to addmm/mm with linear node
Reason is that it simplifies the downstream logic of lowering to just linear node.
Expand All @@ -125,6 +156,9 @@ def replace_addmm_mm_with_linear(graph: torch.fx.Graph) -> torch.fx.Graph:
# Skip this node as it appears to be a standalone `addmm`
continue
weight_node = weight_t_node.args[0]
if not is_constant_tensor(weight_node, exported_program):
# A runtime-computed operand is a matmul, not a linear
continue
args = (node.args[1], weight_node, node.args[0])
linear_node = graph.create_node(
"call_function", ops.aten.linear.default, args
Expand All @@ -142,6 +176,9 @@ def replace_addmm_mm_with_linear(graph: torch.fx.Graph) -> torch.fx.Graph:
# Skip this node as it appears to be a standalone `mm`
continue
weight_node = weight_t_node.args[0]
if not is_constant_tensor(weight_node, exported_program):
# A runtime-computed operand is a matmul, not a linear
continue
args = (node.args[0], weight_node)
linear_node = graph.create_node(
"call_function", ops.aten.linear.default, args
Expand All @@ -158,13 +195,24 @@ def replace_addmm_mm_with_linear(graph: torch.fx.Graph) -> torch.fx.Graph:
return graph


def apply_addmm_mm_to_linear_transform(graph: torch.fx.Graph) -> torch.fx.Graph:
graph = replace_addmm_mm_with_linear(graph)
def apply_addmm_mm_to_linear_transform(
graph: torch.fx.Graph, exported_program: Optional[ExportedProgram] = None
) -> torch.fx.Graph:
graph = replace_addmm_mm_with_linear(graph, exported_program)
graph = replace_linear_view_copy_input_output(graph)
return graph


class AddmmToLinearTransform(ExportPass):
def __init__(self, exported_program: Optional[ExportedProgram] = None) -> None:
super().__init__()
# Backends that run this pass through a pass manager which threads the
# owning program set this attribute themselves; see
# backends/vulkan/vulkan_preprocess.py.
self._exported_program = exported_program

def call(self, graph_module: torch.fx.GraphModule) -> PassResult:
graph_module.graph = apply_addmm_mm_to_linear_transform(graph_module.graph)
graph_module.graph = apply_addmm_mm_to_linear_transform(
graph_module.graph, self._exported_program
)
return PassResult(graph_module, True)
80 changes: 80 additions & 0 deletions backends/transforms/test/test_addmm_mm_to_linear.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

import unittest

import torch
from executorch.backends.transforms.addmm_mm_to_linear import AddmmToLinearTransform
from executorch.exir import to_edge
from executorch.exir.dialects._ops import ops as exir_ops


def count_targets(graph: torch.fx.Graph, target) -> int:
return sum(1 for n in graph.nodes if n.op == "call_function" and n.target == target)


class TestAddmmToLinearTransform(unittest.TestCase):
def _transform(self, model, example_inputs):
edge = to_edge(torch.export.export(model, example_inputs, strict=True))
program = edge.exported_program()
transform = AddmmToLinearTransform(program)
return transform(program.graph_module).graph_module.graph

def test_constant_weight_is_rewritten_to_linear(self):
class Model(torch.nn.Module):
def __init__(self):
super().__init__()
self.fc = torch.nn.Linear(8, 4)

def forward(self, x):
return self.fc(x)

graph = self._transform(Model().eval(), (torch.randn(2, 8),))
self.assertEqual(count_targets(graph, exir_ops.edge.aten.linear.default), 1)
self.assertEqual(count_targets(graph, exir_ops.edge.aten.addmm.default), 0)
self.assertEqual(count_targets(graph, exir_ops.edge.aten.mm.default), 0)

def test_computed_weight_stays_a_matmul(self):
# `w` is produced at runtime, so the transposed matmul below is not a
# linear: backends prepack a linear's weight while building their graph
# and cannot do that for a value that only exists during execution.
class Model(torch.nn.Module):
def forward(self, x, w):
return torch.mm(x, (w * 2.0).t())

graph = self._transform(Model().eval(), (torch.randn(2, 8), torch.randn(4, 8)))
self.assertEqual(count_targets(graph, exir_ops.edge.aten.linear.default), 0)
self.assertEqual(count_targets(graph, exir_ops.edge.aten.mm.default), 1)

def test_computed_bias_operand_stays_an_addmm(self):
class Model(torch.nn.Module):
def forward(self, x, w, b):
return torch.addmm(b, x, (w * 2.0).t())

graph = self._transform(
Model().eval(),
(torch.randn(2, 8), torch.randn(4, 8), torch.randn(4)),
)
self.assertEqual(count_targets(graph, exir_ops.edge.aten.linear.default), 0)
self.assertEqual(count_targets(graph, exir_ops.edge.aten.addmm.default), 1)

def test_user_input_weight_is_not_rewritten(self):
class Model(torch.nn.Module):
def forward(self, x, w):
return torch.nn.functional.linear(x, w)

edge = to_edge(
torch.export.export(
Model().eval(), (torch.randn(2, 8), torch.randn(4, 8)), strict=True
)
)
program = edge.exported_program()
graph = AddmmToLinearTransform(program)(program.graph_module).graph_module.graph
self.assertEqual(count_targets(graph, exir_ops.edge.aten.linear.default), 0)


if __name__ == "__main__":
unittest.main()
Loading