This is an automated email from the ASF dual-hosted git repository.
tqchen pushed a commit to branch unity
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/unity by this push:
new 9daf9acf05 [Unity][Frontend] Some changes on the PyTorch FX Frontend
(#14625)
9daf9acf05 is described below
commit 9daf9acf057e45434411eda62dfa6937c0c983ac
Author: Chaofan Lin <[email protected]>
AuthorDate: Sun Apr 16 02:43:56 2023 +0800
[Unity][Frontend] Some changes on the PyTorch FX Frontend (#14625)
Upstreaming some changes for PyTorch FX Frontend. Add supports for:
torch.nn.functional.cross_entropy
torch.nn.CrossEntropyLoss
torch.nn.AvgPool2d
torch.nn.functional.avg_pool2d
torch.nn.Identity
torch iadd
Co-authored-by: Yixin Dong <[email protected]>
Co-authored-by: Bohan Hou <[email protected]>
---
python/tvm/relax/frontend/torch/fx_translator.py | 69 ++++++++
tests/python/relax/test_frontend_from_fx.py | 203 +++++++++++++++++++++++
2 files changed, 272 insertions(+)
diff --git a/python/tvm/relax/frontend/torch/fx_translator.py
b/python/tvm/relax/frontend/torch/fx_translator.py
index f1ffbeac4b..54890bd3c5 100644
--- a/python/tvm/relax/frontend/torch/fx_translator.py
+++ b/python/tvm/relax/frontend/torch/fx_translator.py
@@ -722,6 +722,34 @@ class TorchFXImporter:
)
)
+ def _avg_pool2d(self, node: fx.node.Node) -> relax.Var:
+ x = self.env[node.args[0]]
+ if node.target in self.named_modules:
+ module = self.named_modules[node.target]
+ kernel = module.kernel_size
+ stride = module.stride
+ padding = module.padding
+ ceil_mode = module.ceil_mode
+ else:
+ nargs = len(node.args)
+ kernel = node.args[1] if nargs > 1 else node.kwargs["kernel_size"]
+ stride = node.args[2] if nargs > 2 else node.kwargs["stride"]
+ padding = node.args[3] if nargs > 3 else node.kwargs["padding"]
+ ceil_mode = node.args[4] if nargs > 4 else node.kwargs["ceil_mode"]
+
+ stride = kernel if stride is None else stride
+
+ return self.block_builder.emit(
+ relax.op.nn.avg_pool2d(
+ x,
+ pool_size=kernel,
+ strides=stride,
+ padding=padding,
+ layout="NCHW",
+ ceil_mode=ceil_mode,
+ )
+ )
+
def _adaptive_avg_pool2d(self, is_module: bool) -> Callable:
from torch import fx
@@ -939,6 +967,41 @@ class TorchFXImporter:
)
)
+ def _cross_entropy(self, node: fx.node.Node) -> relax.Expr:
+ preds = self.env[node.args[0]]
+ targets = self.env[node.args[1]]
+
+ # functional.cross_entropy
+ if node.target not in self.named_modules:
+ weights = node.kwargs["weight"]
+ if weights is not None:
+ weights = self.env[weights]
+ reduction = node.kwargs["reduction"]
+ ignore_index = node.kwargs["ignore_index"]
+
+ return self.block_builder.emit(
+ relax.op.nn.nll_loss(
+ relax.op.nn.log_softmax(preds), targets, weights,
reduction, ignore_index
+ )
+ )
+
+ module = self.named_modules[node.target]
+
+ weights = module.weight
+ if weights is not None:
+ if weights in self.params:
+ weights = self.params[weights]
+ else:
+ weights = relax.const(weights.numpy(), preds.struct_info.dtype)
+ reduction = module.reduction
+ ignore_index = module.ignore_index
+
+ return self.block_builder.emit(
+ relax.op.nn.nll_loss(
+ relax.op.nn.log_softmax(preds), targets, weights, reduction,
ignore_index
+ )
+ )
+
########## Others ##########
def _size(self, node: fx.node.Node) -> relax.Expr:
@@ -1030,6 +1093,7 @@ class TorchFXImporter:
nn.Conv1d: self._conv1d,
nn.Conv2d: self._conv2d,
nn.MaxPool2d: self._max_pool2d,
+ nn.AvgPool2d: self._avg_pool2d,
nn.AdaptiveAvgPool2d: self._adaptive_avg_pool2d(is_module=True),
nn.Softmax: self._softmax,
nn.ReLU: lambda node:
self.block_builder.emit(relax.op.nn.relu(self.env[node.args[0]])),
@@ -1042,11 +1106,14 @@ class TorchFXImporter:
nn.LayerNorm: self._layer_norm,
nn.GroupNorm: self._group_norm,
nn.Dropout: lambda node: self.env[node.args[0]],
+ nn.Identity: lambda node: self.env[node.args[0]],
nn.modules.sparse.Embedding: self._embedding,
+ nn.CrossEntropyLoss: self._cross_entropy,
# call_function and call_method
"cos": self._cos,
"exp": self._exp,
"sin": self._sin,
+ "iadd": self._add,
"add": self._add,
"floordiv": self._floordiv,
"mul": self._mul,
@@ -1105,6 +1172,7 @@ class TorchFXImporter:
"getitem": self._getitem,
"contiguous": lambda node: self.env[node.args[0]],
"to": self._to,
+ "avg_pool2d": self._avg_pool2d,
"adaptive_avg_pool2d": self._adaptive_avg_pool2d(is_module=False),
"layer_norm": self._layer_norm,
"index_select": self._index_select,
@@ -1116,6 +1184,7 @@ class TorchFXImporter:
"rsqrt": self._rsqrt,
"neg": self._neg,
"max": self._max,
+ "cross_entropy": self._cross_entropy,
}
def from_fx(
diff --git a/tests/python/relax/test_frontend_from_fx.py
b/tests/python/relax/test_frontend_from_fx.py
index 2285131e39..4eb7c2afa4 100644
--- a/tests/python/relax/test_frontend_from_fx.py
+++ b/tests/python/relax/test_frontend_from_fx.py
@@ -588,6 +588,83 @@ def test_maxpool2d():
verify_model(MaxPool2d3(), input_info, {}, expected3)
[email protected]_gpu
+def test_avgpool2d():
+ import torch
+ from torch.nn import Module
+
+ torch.set_grad_enabled(False)
+ torch.random.manual_seed(0)
+
+ input_info = [([1, 3, 10, 10], "float32")]
+
+ class AvgPool2d(Module):
+ def __init__(self):
+ super().__init__()
+ self.pool = torch.nn.AvgPool2d(kernel_size=[1, 1])
+
+ def forward(self, input):
+ return self.pool(input)
+
+ @tvm.script.ir_module
+ class expected1:
+ @R.function
+ def main(
+ input_1: R.Tensor((1, 3, 10, 10), dtype="float32")
+ ) -> R.Tensor((1, 3, 10, 10), dtype="float32"):
+ # block 0
+ with R.dataflow():
+ lv: R.Tensor((1, 3, 10, 10), dtype="float32") =
R.nn.avg_pool2d(
+ input_1,
+ pool_size=[1, 1],
+ strides=[1, 1],
+ dilation=[1, 1],
+ padding=[0, 0, 0, 0],
+ layout="NCHW",
+ out_layout="NCHW",
+ )
+ gv: R.Tensor((1, 3, 10, 10), dtype="float32") = lv
+ R.output(gv)
+ return gv
+
+ class AvgPool2d2(Module):
+ def __init__(self):
+ super().__init__()
+ self.pool = torch.nn.AvgPool2d(kernel_size=[4, 4], stride=2,
padding=2, ceil_mode=True)
+
+ def forward(self, input):
+ return self.pool(input)
+
+ class AvgPool2d3(Module):
+ def forward(self, input):
+ return torch.nn.functional.avg_pool2d(
+ input, kernel_size=[4, 4], stride=2, padding=2, ceil_mode=True
+ )
+
+ @tvm.script.ir_module
+ class expected2:
+ @R.function
+ def main(input_1: R.Tensor((1, 3, 10, 10), dtype="float32")):
+ with R.dataflow():
+ lv = R.nn.avg_pool2d(
+ input_1,
+ pool_size=[4, 4],
+ strides=[2, 2],
+ dilation=[1, 1],
+ padding=[2, 2, 2, 2],
+ ceil_mode=True,
+ layout="NCHW",
+ out_layout="NCHW",
+ )
+ gv = lv
+ R.output(gv)
+ return gv
+
+ verify_model(AvgPool2d(), input_info, {}, expected1)
+ verify_model(AvgPool2d2(), input_info, {}, expected2)
+ verify_model(AvgPool2d3(), input_info, {}, expected2)
+
+
@tvm.testing.requires_gpu
def test_adaptive_avgpool2d():
import torch
@@ -902,6 +979,132 @@ def test_functional_layernorm():
verify_model(model, input_info, binding, expected1)
[email protected]_gpu
+def test_cross_entropy():
+ import torch
+ from torch.nn import Module
+
+ torch.set_grad_enabled(False)
+ torch.random.manual_seed(0)
+
+ input_info = [([3, 2], "float32"), ([3], "int32")]
+
+ class CrossEntropy1(Module):
+ def __init__(self):
+ super().__init__()
+ self.loss = torch.nn.CrossEntropyLoss()
+
+ def forward(self, logits, targets):
+ return self.loss(logits, targets)
+
+ @tvm.script.ir_module
+ class expected1:
+ @R.function
+ def main(
+ inp_0: R.Tensor((3, 2), dtype="float32"), inp_1: R.Tensor((3,),
dtype="int32")
+ ) -> R.Tensor((), dtype="float32"):
+ with R.dataflow():
+ lv: R.Tensor((3, 2), dtype="float32") =
R.nn.log_softmax(inp_0, axis=-1)
+ lv1: R.Tensor((), dtype="float32") = R.nn.nll_loss(
+ lv, inp_1, reduction="mean", ignore_index=-100
+ )
+ gv: R.Tensor((), dtype="float32") = lv1
+ R.output(gv)
+ return gv
+
+ class CrossEntropy2(Module):
+ def __init__(self):
+ super().__init__()
+ self.weight = torch.nn.Parameter(torch.ones((2,)))
+ self.loss = torch.nn.CrossEntropyLoss(weight=self.weight)
+
+ def forward(self, logits, targets):
+ return self.loss(logits, targets)
+
+ @tvm.script.ir_module
+ class expected2:
+ @R.function
+ def main(
+ inp_0: R.Tensor((3, 2), dtype="float32"),
+ inp_1: R.Tensor((3,), dtype="int32"),
+ w1: R.Tensor((2,), dtype="float32"),
+ ) -> R.Tensor((), dtype="float32"):
+ with R.dataflow():
+ lv: R.Tensor((3, 2), dtype="float32") =
R.nn.log_softmax(inp_0, axis=-1)
+ lv1: R.Tensor((), dtype="float32") = R.nn.nll_loss(
+ lv,
+ inp_1,
+ w1,
+ reduction="mean",
+ ignore_index=-100,
+ )
+ gv: R.Tensor((), dtype="float32") = lv1
+ R.output(gv)
+ return gv
+
+ class CrossEntropy3(Module):
+ def __init__(self):
+ super().__init__()
+ self.loss = torch.nn.CrossEntropyLoss(ignore_index=1,
reduction="sum")
+
+ def forward(self, logits, targets):
+ return self.loss(logits, targets)
+
+ @tvm.script.ir_module
+ class expected3:
+ @R.function
+ def main(
+ inp_0: R.Tensor((3, 2), dtype="float32"), inp_1: R.Tensor((3,),
dtype="int32")
+ ) -> R.Tensor((), dtype="float32"):
+ with R.dataflow():
+ lv: R.Tensor((3, 2), dtype="float32") =
R.nn.log_softmax(inp_0, axis=-1)
+ lv1: R.Tensor((), dtype="float32") = R.nn.nll_loss(
+ lv, inp_1, reduction="sum", ignore_index=1
+ )
+ gv: R.Tensor((), dtype="float32") = lv1
+ R.output(gv)
+ return gv
+
+ verify_model(CrossEntropy1(), input_info, {}, expected1)
+ model = CrossEntropy2()
+ binding = {"w1": model.loss.weight.numpy()}
+ verify_model(model, input_info, binding, expected2)
+ verify_model(CrossEntropy3(), input_info, {}, expected3)
+
+
[email protected]_gpu
+def test_functional_cross_entropy():
+ import torch
+ from torch.nn import Module
+
+ torch.set_grad_enabled(False)
+ torch.random.manual_seed(0)
+
+ input_info = [([3, 10], "float32"), ([3], "int32")]
+
+ class CrossEntropy(Module):
+ def forward(self, logits, targets):
+ return torch.nn.functional.cross_entropy(logits, targets)
+
+ @tvm.script.ir_module
+ class expected1:
+ @R.function
+ def main(
+ inp_0: R.Tensor((3, 10), dtype="float32"), inp_1: R.Tensor((3,),
dtype="int32")
+ ) -> R.Tensor((), dtype="float32"):
+ with R.dataflow():
+ lv: R.Tensor((3, 10), dtype="float32") =
R.nn.log_softmax(inp_0, axis=-1)
+ lv1: R.Tensor((), dtype="float32") = R.nn.nll_loss(
+ lv, inp_1, reduction="mean", ignore_index=-100
+ )
+ gv: R.Tensor((), dtype="float32") = lv1
+ R.output(gv)
+ return gv
+
+ model = CrossEntropy()
+ verify_model(model, input_info, {}, expected1)
+
+
@tvm.testing.requires_gpu
def test_silu():
import torch