{
  "id": 388619,
  "title": "The loss function in the Graphnet paper (vMF)",
  "url": "/competitions/icecube-neutrinos-in-deep-ice/discussion/388619",
  "author_name": "",
  "post_date": "2023-02-18T15:43:43.250676700Z",
  "votes": 21,
  "comment_count": 1,
  "views": 0,
  "content": "<p>The von Mises-Fisher (vMF) loss is a loss function used in machine learning for training neural networks that output high-dimensional vectors on the hypersphere. It is particularly useful in tasks such as clustering, classification, and representation learning of angular data, such as images, text, or audio signals.</p>\n<p>The vMF distribution is a probability distribution defined on the hypersphere, which is a higher-dimensional generalization of the circle. It is often used to model directional data and is characterized by two parameters: the mean direction, which is a unit vector on the hypersphere, and the concentration parameter, which determines the spread of the distribution.</p>\n<p>The vMF loss measures the dissimilarity between two unit vectors on the hypersphere, which can be seen as points on the surface of the sphere. It is defined as the negative log-likelihood of the predicted vector given the ground-truth vector under the vMF distribution with a fixed concentration parameter. The vMF loss encourages the predicted vector to be close to the ground-truth vector in terms of their angular distance and penalizes large deviations from the mean direction.</p>\n<p>The vMF loss can be computed efficiently using standard vector operations and is differentiable, making it suitable for gradient-based optimization algorithms such as stochastic gradient descent.</p>\n<p>Here is the code snippet from <a href=\"https://github.com/graphnet-team/graphnet/blob/b77b25c7517b6a81203e300a79258810701f815c/src/graphnet/training/loss_functions.py\" target=\"_blank\">Graphnet</a> repository:</p>\n<pre><code> ():\n    \n\n     () -&gt; Tensor:\n        \n        target = target.reshape(-, )\n        \n         prediction.dim() ==   prediction.size()[] == \n         target.dim() == \n         prediction.size()[] == target.size()[]\n\n        kappa = prediction[:, ]\n        p = kappa.unsqueeze() * prediction[:, [, , ]]\n         self._evaluate(p, target)\n</code></pre>\n<p>This is a Python class named VonMisesFisher3DLoss that inherits from a parent class VonMisesFisherLoss and provides a specific implementation of the von Mises-Fisher loss function for 3D vectors.</p>\n<p>The purpose of the class is to define a method _forward() that calculates the von Mises-Fisher loss between the predicted 3D vectors and the ground-truth 3D vectors, given a parameter kappa that represents the concentration parameter of the vMF distribution. The method takes two arguments: prediction and target. The prediction tensor is assumed to have shape [N, 4], where the first three columns correspond to the predicted directions and the last column corresponds to the predicted value of kappa. The target tensor is assumed to have shape [N, 3], representing the ground-truth directions.</p>\n<p>The _forward() method first reshapes the target tensor to have shape [N, 3]. It then performs some input validation to ensure that the dimensions of prediction and target are compatible. The method then extracts the predicted kappa values and multiplies them with the predicted directions to obtain a weighted prediction p. Finally, the method calls a helper method _evaluate() to compute the elementwise von Mises-Fisher loss terms between p and target and returns the resulting tensor of shape [N,].</p>\n<p>Overall, this class provides an implementation of the von Mises-Fisher loss function that is tailored for 3D vectors, which can be used to train a neural network to predict 3D directions or to evaluate the performance of a model that outputs such directions.</p>\n<p>Here is the parent class:</p>\n<pre><code> ():\n    \n\n\n     () -&gt; Tensor:  \n        \n         LogCMK.apply(m, kappa)\n\n\n     () -&gt; Tensor:  \n        \n        v = m /  - \n        a = torch.sqrt((v + ) **  + kappa**)\n        b = v - \n         -a + b * torch.log(b + a)\n\n\n     () -&gt; Tensor:  \n        \n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa &lt; kappa_switch\n\n        \n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n         ret\n\n     () -&gt; Tensor:\n        \n        \n         prediction.dim() == \n         target.dim() == \n         prediction.size() == target.size()\n\n        \n        m = target.size()[]\n        k = torch.norm(prediction, dim=)\n        dotprod = torch.(prediction * target, dim=)\n        elements = -self.log_cmk(m, k) - dotprod\n         elements\n\n\n     () -&gt; Tensor:\n         NotImplementedError\n</code></pre>\n<p>This is a Python class named VonMisesFisherLoss that defines a general interface for calculating the von Mises-Fisher loss function for a vector of arbitrary dimension D. The class is an abstract base class that inherits from a parent class LossFunction.</p>\n<p>The class defines three class methods to compute the log_cmk term of the von Mises-Fisher loss function, which corresponds to the logarithm of the normalizing constant of the von Mises-Fisher distribution. The log_cmk_exact() method computes the exact value of log_cmk for a given dimension m and concentration parameter kappa, using a custom autograd function LogCMK.apply() that computes the value of the modified Bessel function of the second kind. The log_cmk_approx() method computes an approximate value of log_cmk for large values of kappa using a closed-form expression from a research paper. The log_cmk() method automatically switches between the exact and approximate methods depending on the value of kappa and ensures continuity at a given threshold kappa_switch.</p>\n<p>The class also defines a private method _evaluate() that calculates the von Mises-Fisher loss between a predicted vector of shape [batch_size, D] and a target unit vector of the same shape. The method uses the torch.norm() function to compute the Euclidean norm of the predicted vector, the torch.sum() function to compute the dot product between the predicted and target vectors, and the log_cmk() method to compute the log_cmk term. The method returns a tensor of shape [batch_size,] that contains the elementwise von Mises-Fisher loss terms.</p>\n<p>Finally, the class defines an abstract method _forward() that must be implemented by subclasses to compute the von Mises-Fisher loss for a vector of a specific dimension. This method is not implemented in the base class because the implementation of the von Mises-Fisher loss function depends on the dimensionality of the vectors, and a specific implementation is required for each dimension. Therefore, the class serves as a template for defining von Mises-Fisher loss functions for different vector dimensions.</p>",
  "messages": [
    {
      "id": "2149716",
      "postDate": "02/18/2023 15:43:43",
      "content": "<p>The von Mises-Fisher (vMF) loss is a loss function used in machine learning for training neural networks that output high-dimensional vectors on the hypersphere. It is particularly useful in tasks such as clustering, classification, and representation learning of angular data, such as images, text, or audio signals.</p>\n<p>The vMF distribution is a probability distribution defined on the hypersphere, which is a higher-dimensional generalization of the circle. It is often used to model directional data and is characterized by two parameters: the mean direction, which is a unit vector on the hypersphere, and the concentration parameter, which determines the spread of the distribution.</p>\n<p>The vMF loss measures the dissimilarity between two unit vectors on the hypersphere, which can be seen as points on the surface of the sphere. It is defined as the negative log-likelihood of the predicted vector given the ground-truth vector under the vMF distribution with a fixed concentration parameter. The vMF loss encourages the predicted vector to be close to the ground-truth vector in terms of their angular distance and penalizes large deviations from the mean direction.</p>\n<p>The vMF loss can be computed efficiently using standard vector operations and is differentiable, making it suitable for gradient-based optimization algorithms such as stochastic gradient descent.</p>\n<p>Here is the code snippet from <a href=\"https://github.com/graphnet-team/graphnet/blob/b77b25c7517b6a81203e300a79258810701f815c/src/graphnet/training/loss_functions.py\" target=\"_blank\">Graphnet</a> repository:</p>\n<pre><code> ():\n    \n\n     () -&gt; Tensor:\n        \n        target = target.reshape(-, )\n        \n         prediction.dim() ==   prediction.size()[] == \n         target.dim() == \n         prediction.size()[] == target.size()[]\n\n        kappa = prediction[:, ]\n        p = kappa.unsqueeze() * prediction[:, [, , ]]\n         self._evaluate(p, target)\n</code></pre>\n<p>This is a Python class named VonMisesFisher3DLoss that inherits from a parent class VonMisesFisherLoss and provides a specific implementation of the von Mises-Fisher loss function for 3D vectors.</p>\n<p>The purpose of the class is to define a method _forward() that calculates the von Mises-Fisher loss between the predicted 3D vectors and the ground-truth 3D vectors, given a parameter kappa that represents the concentration parameter of the vMF distribution. The method takes two arguments: prediction and target. The prediction tensor is assumed to have shape [N, 4], where the first three columns correspond to the predicted directions and the last column corresponds to the predicted value of kappa. The target tensor is assumed to have shape [N, 3], representing the ground-truth directions.</p>\n<p>The _forward() method first reshapes the target tensor to have shape [N, 3]. It then performs some input validation to ensure that the dimensions of prediction and target are compatible. The method then extracts the predicted kappa values and multiplies them with the predicted directions to obtain a weighted prediction p. Finally, the method calls a helper method _evaluate() to compute the elementwise von Mises-Fisher loss terms between p and target and returns the resulting tensor of shape [N,].</p>\n<p>Overall, this class provides an implementation of the von Mises-Fisher loss function that is tailored for 3D vectors, which can be used to train a neural network to predict 3D directions or to evaluate the performance of a model that outputs such directions.</p>\n<p>Here is the parent class:</p>\n<pre><code> ():\n    \n\n\n     () -&gt; Tensor:  \n        \n         LogCMK.apply(m, kappa)\n\n\n     () -&gt; Tensor:  \n        \n        v = m /  - \n        a = torch.sqrt((v + ) **  + kappa**)\n        b = v - \n         -a + b * torch.log(b + a)\n\n\n     () -&gt; Tensor:  \n        \n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa &lt; kappa_switch\n\n        \n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n         ret\n\n     () -&gt; Tensor:\n        \n        \n         prediction.dim() == \n         target.dim() == \n         prediction.size() == target.size()\n\n        \n        m = target.size()[]\n        k = torch.norm(prediction, dim=)\n        dotprod = torch.(prediction * target, dim=)\n        elements = -self.log_cmk(m, k) - dotprod\n         elements\n\n\n     () -&gt; Tensor:\n         NotImplementedError\n</code></pre>\n<p>This is a Python class named VonMisesFisherLoss that defines a general interface for calculating the von Mises-Fisher loss function for a vector of arbitrary dimension D. The class is an abstract base class that inherits from a parent class LossFunction.</p>\n<p>The class defines three class methods to compute the log_cmk term of the von Mises-Fisher loss function, which corresponds to the logarithm of the normalizing constant of the von Mises-Fisher distribution. The log_cmk_exact() method computes the exact value of log_cmk for a given dimension m and concentration parameter kappa, using a custom autograd function LogCMK.apply() that computes the value of the modified Bessel function of the second kind. The log_cmk_approx() method computes an approximate value of log_cmk for large values of kappa using a closed-form expression from a research paper. The log_cmk() method automatically switches between the exact and approximate methods depending on the value of kappa and ensures continuity at a given threshold kappa_switch.</p>\n<p>The class also defines a private method _evaluate() that calculates the von Mises-Fisher loss between a predicted vector of shape [batch_size, D] and a target unit vector of the same shape. The method uses the torch.norm() function to compute the Euclidean norm of the predicted vector, the torch.sum() function to compute the dot product between the predicted and target vectors, and the log_cmk() method to compute the log_cmk term. The method returns a tensor of shape [batch_size,] that contains the elementwise von Mises-Fisher loss terms.</p>\n<p>Finally, the class defines an abstract method _forward() that must be implemented by subclasses to compute the von Mises-Fisher loss for a vector of a specific dimension. This method is not implemented in the base class because the implementation of the von Mises-Fisher loss function depends on the dimensionality of the vectors, and a specific implementation is required for each dimension. Therefore, the class serves as a template for defining von Mises-Fisher loss functions for different vector dimensions.</p>",
      "rawMarkdown": "The von Mises-Fisher (vMF) loss is a loss function used in machine learning for training neural networks that output high-dimensional vectors on the hypersphere. It is particularly useful in tasks such as clustering, classification, and representation learning of angular data, such as images, text, or audio signals.\n\nThe vMF distribution is a probability distribution defined on the hypersphere, which is a higher-dimensional generalization of the circle. It is often used to model directional data and is characterized by two parameters: the mean direction, which is a unit vector on the hypersphere, and the concentration parameter, which determines the spread of the distribution.\n\nThe vMF loss measures the dissimilarity between two unit vectors on the hypersphere, which can be seen as points on the surface of the sphere. It is defined as the negative log-likelihood of the predicted vector given the ground-truth vector under the vMF distribution with a fixed concentration parameter. The vMF loss encourages the predicted vector to be close to the ground-truth vector in terms of their angular distance and penalizes large deviations from the mean direction.\n\nThe vMF loss can be computed efficiently using standard vector operations and is differentiable, making it suitable for gradient-based optimization algorithms such as stochastic gradient descent.\n\nHere is the code snippet from [Graphnet](https://github.com/graphnet-team/graphnet/blob/b77b25c7517b6a81203e300a79258810701f815c/src/graphnet/training/loss_functions.py) repository:\n\n```python\nclass VonMisesFisher3DLoss(VonMisesFisherLoss):\n    \"\"\"von Mises-Fisher loss function vectors in the 3D plane.\"\"\"\n\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a direction in the 3D.\n        Args:\n            prediction: Output of the model. Must have shape [N, 4] where\n                columns 0, 1, 2 are predictions of `direction` and last column\n                is an estimate of `kappa`.\n            target: Target tensor, extracted from graph object.\n        Returns:\n            Elementwise von Mises-Fisher loss terms. Shape [N,]\n        \"\"\"\n        target = target.reshape(-1, 3)\n        # Check(s)\n        assert prediction.dim() == 2 and prediction.size()[1] == 4\n        assert target.dim() == 2\n        assert prediction.size()[0] == target.size()[0]\n\n        kappa = prediction[:, 3]\n        p = kappa.unsqueeze(1) * prediction[:, [0, 1, 2]]\n        return self._evaluate(p, target)\n\n```\n\nThis is a Python class named VonMisesFisher3DLoss that inherits from a parent class VonMisesFisherLoss and provides a specific implementation of the von Mises-Fisher loss function for 3D vectors.\n\nThe purpose of the class is to define a method _forward() that calculates the von Mises-Fisher loss between the predicted 3D vectors and the ground-truth 3D vectors, given a parameter kappa that represents the concentration parameter of the vMF distribution. The method takes two arguments: prediction and target. The prediction tensor is assumed to have shape [N, 4], where the first three columns correspond to the predicted directions and the last column corresponds to the predicted value of kappa. The target tensor is assumed to have shape [N, 3], representing the ground-truth directions.\n\nThe _forward() method first reshapes the target tensor to have shape [N, 3]. It then performs some input validation to ensure that the dimensions of prediction and target are compatible. The method then extracts the predicted kappa values and multiplies them with the predicted directions to obtain a weighted prediction p. Finally, the method calls a helper method _evaluate() to compute the elementwise von Mises-Fisher loss terms between p and target and returns the resulting tensor of shape [N,].\n\nOverall, this class provides an implementation of the von Mises-Fisher loss function that is tailored for 3D vectors, which can be used to train a neural network to predict 3D directions or to evaluate the performance of a model that outputs such directions.\n\nHere is the parent class:\n\n```python\nclass VonMisesFisherLoss(LossFunction):\n    \"\"\"General class for calculating von Mises-Fisher loss.\n    Requires implementation for specific dimension `m` in which the target and\n    prediction vectors need to be prepared.\n    \"\"\"\n\n    @classmethod\n    def log_cmk_exact(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss exactly.\"\"\"\n        return LogCMK.apply(m, kappa)\n\n    @classmethod\n    def log_cmk_approx(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss approx.\n        [https://arxiv.org/abs/1812.04616] Sec. 8.2 with additional minus sign.\n        \"\"\"\n        v = m / 2.0 - 0.5\n        a = torch.sqrt((v + 1) ** 2 + kappa**2)\n        b = v - 1\n        return -a + b * torch.log(b + a)\n\n    @classmethod\n    def log_cmk(\n        cls, m: int, kappa: Tensor, kappa_switch: float = 100.0\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss.\n        Since `log_cmk_exact` is diverges for `kappa` >~ 700 (using float64\n        precision), and since `log_cmk_approx` is unaccurate for small `kappa`,\n        this method automatically switches between the two at `kappa_switch`,\n        ensuring continuity at this point.\n        \"\"\"\n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa < kappa_switch\n\n        # Ensure continuity at `kappa_switch`\n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n        return ret\n\n    def _evaluate(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a vector in D dimensons.\n        This loss utilises the von Mises-Fisher distribution, which is a\n        probability distribution on the (D - 1) sphere in D-dimensional space.\n        Args:\n            prediction: Predicted vector, of shape [batch_size, D].\n            target: Target unit vector, of shape [batch_size, D].\n        Returns:\n            Elementwise von Mises-Fisher loss terms.\n        \"\"\"\n        # Check(s)\n        assert prediction.dim() == 2\n        assert target.dim() == 2\n        assert prediction.size() == target.size()\n\n        # Computing loss\n        m = target.size()[1]\n        k = torch.norm(prediction, dim=1)\n        dotprod = torch.sum(prediction * target, dim=1)\n        elements = -self.log_cmk(m, k) - dotprod\n        return elements\n\n    @abstractmethod\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        raise NotImplementedError\n```\n\nThis is a Python class named VonMisesFisherLoss that defines a general interface for calculating the von Mises-Fisher loss function for a vector of arbitrary dimension D. The class is an abstract base class that inherits from a parent class LossFunction.\n\nThe class defines three class methods to compute the log_cmk term of the von Mises-Fisher loss function, which corresponds to the logarithm of the normalizing constant of the von Mises-Fisher distribution. The log_cmk_exact() method computes the exact value of log_cmk for a given dimension m and concentration parameter kappa, using a custom autograd function LogCMK.apply() that computes the value of the modified Bessel function of the second kind. The log_cmk_approx() method computes an approximate value of log_cmk for large values of kappa using a closed-form expression from a research paper. The log_cmk() method automatically switches between the exact and approximate methods depending on the value of kappa and ensures continuity at a given threshold kappa_switch.\n\nThe class also defines a private method _evaluate() that calculates the von Mises-Fisher loss between a predicted vector of shape [batch_size, D] and a target unit vector of the same shape. The method uses the torch.norm() function to compute the Euclidean norm of the predicted vector, the torch.sum() function to compute the dot product between the predicted and target vectors, and the log_cmk() method to compute the log_cmk term. The method returns a tensor of shape [batch_size,] that contains the elementwise von Mises-Fisher loss terms.\n\nFinally, the class defines an abstract method _forward() that must be implemented by subclasses to compute the von Mises-Fisher loss for a vector of a specific dimension. This method is not implemented in the base class because the implementation of the von Mises-Fisher loss function depends on the dimensionality of the vectors, and a specific implementation is required for each dimension. Therefore, the class serves as a template for defining von Mises-Fisher loss functions for different vector dimensions.",
      "votes": null
    },
    {
      "id": "2150381",
      "postDate": "02/19/2023 07:44:48",
      "content": "<p>Good comment on the code!<br>\nIt made the code more readable!</p>\n<p>As a newbie, it's first to see this vMF loss function.<br>\nIs there any human-readable <strong>interpretation</strong> of the loss number?<br>\nFor example… <em>the vMF loss 2.03 means your mean-angular-distance is about 2.03 / 2 rad</em>…</p>",
      "rawMarkdown": "Good comment on the code!\nIt made the code more readable!\n\nAs a newbie, it's first to see this vMF loss function.\nIs there any human-readable **interpretation** of the loss number?\nFor example... *the vMF loss 2.03 means your mean-angular-distance is about 2.03 / 2 rad*...",
      "votes": null
    }
  ],
  "comments": [
    {
      "id": 2150381,
      "author_name": "seungmoklee",
      "author_url": "",
      "post_date": "02/19/2023 07:44:48",
      "content": "<p>Good comment on the code!<br>\nIt made the code more readable!</p>\n<p>As a newbie, it's first to see this vMF loss function.<br>\nIs there any human-readable <strong>interpretation</strong> of the loss number?<br>\nFor example… <em>the vMF loss 2.03 means your mean-angular-distance is about 2.03 / 2 rad</em>…</p>",
      "votes": null,
      "replies": []
    }
  ],
  "raw_markdown_by_id": {
    "2149716": "The von Mises-Fisher (vMF) loss is a loss function used in machine learning for training neural networks that output high-dimensional vectors on the hypersphere. It is particularly useful in tasks such as clustering, classification, and representation learning of angular data, such as images, text, or audio signals.\n\nThe vMF distribution is a probability distribution defined on the hypersphere, which is a higher-dimensional generalization of the circle. It is often used to model directional data and is characterized by two parameters: the mean direction, which is a unit vector on the hypersphere, and the concentration parameter, which determines the spread of the distribution.\n\nThe vMF loss measures the dissimilarity between two unit vectors on the hypersphere, which can be seen as points on the surface of the sphere. It is defined as the negative log-likelihood of the predicted vector given the ground-truth vector under the vMF distribution with a fixed concentration parameter. The vMF loss encourages the predicted vector to be close to the ground-truth vector in terms of their angular distance and penalizes large deviations from the mean direction.\n\nThe vMF loss can be computed efficiently using standard vector operations and is differentiable, making it suitable for gradient-based optimization algorithms such as stochastic gradient descent.\n\nHere is the code snippet from [Graphnet](https://github.com/graphnet-team/graphnet/blob/b77b25c7517b6a81203e300a79258810701f815c/src/graphnet/training/loss_functions.py) repository:\n\n```python\nclass VonMisesFisher3DLoss(VonMisesFisherLoss):\n    \"\"\"von Mises-Fisher loss function vectors in the 3D plane.\"\"\"\n\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a direction in the 3D.\n        Args:\n            prediction: Output of the model. Must have shape [N, 4] where\n                columns 0, 1, 2 are predictions of `direction` and last column\n                is an estimate of `kappa`.\n            target: Target tensor, extracted from graph object.\n        Returns:\n            Elementwise von Mises-Fisher loss terms. Shape [N,]\n        \"\"\"\n        target = target.reshape(-1, 3)\n        # Check(s)\n        assert prediction.dim() == 2 and prediction.size()[1] == 4\n        assert target.dim() == 2\n        assert prediction.size()[0] == target.size()[0]\n\n        kappa = prediction[:, 3]\n        p = kappa.unsqueeze(1) * prediction[:, [0, 1, 2]]\n        return self._evaluate(p, target)\n\n```\n\nThis is a Python class named VonMisesFisher3DLoss that inherits from a parent class VonMisesFisherLoss and provides a specific implementation of the von Mises-Fisher loss function for 3D vectors.\n\nThe purpose of the class is to define a method _forward() that calculates the von Mises-Fisher loss between the predicted 3D vectors and the ground-truth 3D vectors, given a parameter kappa that represents the concentration parameter of the vMF distribution. The method takes two arguments: prediction and target. The prediction tensor is assumed to have shape [N, 4], where the first three columns correspond to the predicted directions and the last column corresponds to the predicted value of kappa. The target tensor is assumed to have shape [N, 3], representing the ground-truth directions.\n\nThe _forward() method first reshapes the target tensor to have shape [N, 3]. It then performs some input validation to ensure that the dimensions of prediction and target are compatible. The method then extracts the predicted kappa values and multiplies them with the predicted directions to obtain a weighted prediction p. Finally, the method calls a helper method _evaluate() to compute the elementwise von Mises-Fisher loss terms between p and target and returns the resulting tensor of shape [N,].\n\nOverall, this class provides an implementation of the von Mises-Fisher loss function that is tailored for 3D vectors, which can be used to train a neural network to predict 3D directions or to evaluate the performance of a model that outputs such directions.\n\nHere is the parent class:\n\n```python\nclass VonMisesFisherLoss(LossFunction):\n    \"\"\"General class for calculating von Mises-Fisher loss.\n    Requires implementation for specific dimension `m` in which the target and\n    prediction vectors need to be prepared.\n    \"\"\"\n\n    @classmethod\n    def log_cmk_exact(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss exactly.\"\"\"\n        return LogCMK.apply(m, kappa)\n\n    @classmethod\n    def log_cmk_approx(\n        cls, m: int, kappa: Tensor\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss approx.\n        [https://arxiv.org/abs/1812.04616] Sec. 8.2 with additional minus sign.\n        \"\"\"\n        v = m / 2.0 - 0.5\n        a = torch.sqrt((v + 1) ** 2 + kappa**2)\n        b = v - 1\n        return -a + b * torch.log(b + a)\n\n    @classmethod\n    def log_cmk(\n        cls, m: int, kappa: Tensor, kappa_switch: float = 100.0\n    ) -> Tensor:  # pylint: disable=invalid-name\n        \"\"\"Calculate $log C_{m}(k)$ term in von Mises-Fisher loss.\n        Since `log_cmk_exact` is diverges for `kappa` >~ 700 (using float64\n        precision), and since `log_cmk_approx` is unaccurate for small `kappa`,\n        this method automatically switches between the two at `kappa_switch`,\n        ensuring continuity at this point.\n        \"\"\"\n        kappa_switch = torch.tensor([kappa_switch]).to(kappa.device)\n        mask_exact = kappa < kappa_switch\n\n        # Ensure continuity at `kappa_switch`\n        offset = cls.log_cmk_approx(m, kappa_switch) - cls.log_cmk_exact(\n            m, kappa_switch\n        )\n        ret = cls.log_cmk_approx(m, kappa) - offset\n        ret[mask_exact] = cls.log_cmk_exact(m, kappa[mask_exact])\n        return ret\n\n    def _evaluate(self, prediction: Tensor, target: Tensor) -> Tensor:\n        \"\"\"Calculate von Mises-Fisher loss for a vector in D dimensons.\n        This loss utilises the von Mises-Fisher distribution, which is a\n        probability distribution on the (D - 1) sphere in D-dimensional space.\n        Args:\n            prediction: Predicted vector, of shape [batch_size, D].\n            target: Target unit vector, of shape [batch_size, D].\n        Returns:\n            Elementwise von Mises-Fisher loss terms.\n        \"\"\"\n        # Check(s)\n        assert prediction.dim() == 2\n        assert target.dim() == 2\n        assert prediction.size() == target.size()\n\n        # Computing loss\n        m = target.size()[1]\n        k = torch.norm(prediction, dim=1)\n        dotprod = torch.sum(prediction * target, dim=1)\n        elements = -self.log_cmk(m, k) - dotprod\n        return elements\n\n    @abstractmethod\n    def _forward(self, prediction: Tensor, target: Tensor) -> Tensor:\n        raise NotImplementedError\n```\n\nThis is a Python class named VonMisesFisherLoss that defines a general interface for calculating the von Mises-Fisher loss function for a vector of arbitrary dimension D. The class is an abstract base class that inherits from a parent class LossFunction.\n\nThe class defines three class methods to compute the log_cmk term of the von Mises-Fisher loss function, which corresponds to the logarithm of the normalizing constant of the von Mises-Fisher distribution. The log_cmk_exact() method computes the exact value of log_cmk for a given dimension m and concentration parameter kappa, using a custom autograd function LogCMK.apply() that computes the value of the modified Bessel function of the second kind. The log_cmk_approx() method computes an approximate value of log_cmk for large values of kappa using a closed-form expression from a research paper. The log_cmk() method automatically switches between the exact and approximate methods depending on the value of kappa and ensures continuity at a given threshold kappa_switch.\n\nThe class also defines a private method _evaluate() that calculates the von Mises-Fisher loss between a predicted vector of shape [batch_size, D] and a target unit vector of the same shape. The method uses the torch.norm() function to compute the Euclidean norm of the predicted vector, the torch.sum() function to compute the dot product between the predicted and target vectors, and the log_cmk() method to compute the log_cmk term. The method returns a tensor of shape [batch_size,] that contains the elementwise von Mises-Fisher loss terms.\n\nFinally, the class defines an abstract method _forward() that must be implemented by subclasses to compute the von Mises-Fisher loss for a vector of a specific dimension. This method is not implemented in the base class because the implementation of the von Mises-Fisher loss function depends on the dimensionality of the vectors, and a specific implementation is required for each dimension. Therefore, the class serves as a template for defining von Mises-Fisher loss functions for different vector dimensions.",
    "2150381": "Good comment on the code!\nIt made the code more readable!\n\nAs a newbie, it's first to see this vMF loss function.\nIs there any human-readable **interpretation** of the loss number?\nFor example... *the vMF loss 2.03 means your mean-angular-distance is about 2.03 / 2 rad*..."
  },
  "source": "meta"
}