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

Reply via email to