From d3dc332c2e01d9eee11558741ffe44381eea3c9f Mon Sep 17 00:00:00 2001 From: chenyu Date: Thu, 9 May 2024 20:49:06 -0400 Subject: [PATCH] Tensor.logsumexp (#4442) the subtract max part should share with safe softmax cleaner --- test/test_ops.py | 8 ++++++++ tinygrad/tensor.py | 4 ++++ 2 files changed, 12 insertions(+) diff --git a/test/test_ops.py b/test/test_ops.py index 9092b84bb7..3c89471ce0 100644 --- a/test/test_ops.py +++ b/test/test_ops.py @@ -810,6 +810,14 @@ class TestOps(unittest.TestCase): helper_test_op([(10,10,10)], lambda x: x.log_softmax(1), atol=1e-7, grad_atol=1e-7) helper_test_op([(10,10,10)], lambda x: x.log_softmax(2), atol=1e-7, grad_atol=1e-7) + def test_logsumexp(self): + helper_test_op([(45,65)], lambda x: torch.logsumexp(x, dim=0), lambda x: x.logsumexp(0), atol=1e-7, grad_atol=1e-7) + helper_test_op([(45,65)], lambda x: torch.logsumexp(x, dim=0, keepdim=True), lambda x: x.logsumexp(0, True), atol=1e-7, grad_atol=1e-7) + helper_test_op([(45,65)], lambda x: torch.logsumexp(x, dim=1), lambda x: x.logsumexp(1), atol=1e-7, grad_atol=1e-7) + helper_test_op([(45)], lambda x: torch.logsumexp(x, dim=0), lambda x: x.logsumexp(0), atol=1e-7, grad_atol=1e-7) + helper_test_op([()], lambda x: torch.logsumexp(x, dim=0), lambda x: x.logsumexp(0), atol=1e-7, grad_atol=1e-7) + helper_test_op([()], lambda x: torch.logsumexp(x, dim=-1), lambda x: x.logsumexp(-1), atol=1e-7, grad_atol=1e-7) + def test_sinh(self): helper_test_op([(45,65)], lambda x: x.sinh(), grad_atol=1e-6) # TODO: backward nan instead of inf diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index fb724f6899..3a062e3942 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -951,6 +951,10 @@ class Tensor: m, _, ss = self._softmax(axis) return m - ss.log() + def logsumexp(self, axis=None, keepdim=False): + m = self.max(axis=axis, keepdim=True) + return (self - m).exp().sum(axis=axis, keepdim=keepdim).log() + m.squeeze(axis) + def argmax(self, axis=None, keepdim=False): # NOTE: return the first index if there are multiple occurrences of the maximum values if axis is None: