{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Following the paper https://arxiv.org/abs/2103.15718, I implemented the von Mises–Fisher Loss for three-dimensional vectors.\n\nFrom the paper, for an $n$-dimensional unit vector $\\boldsymbol{x}$ is\n\n$$\np(\\boldsymbol{x} ; \\boldsymbol{\\mu}, \\kappa)=C_n(\\kappa) \\exp \\left(\\kappa \\boldsymbol{\\mu}^{\\mathrm{T}} \\boldsymbol{x}\\right),\n$$\nwhere $C_n(\\kappa)$ is defined as\n$$\nC_n(\\kappa)=\\frac{\\kappa^{n / 2-1}}{(2 \\pi)^{n / 2} I_{n / 2-1}(\\kappa)}.\n$$\n$I_{v}(x)$ is Bessel function of the first kind at order $v$.\n\nFor $n=3$, the $C_n(\\kappa)$ reduces to\n\n$$\nC_3(\\kappa)=\\frac{\\kappa^{1 / 2}}{(2 \\pi)^{3/2} I_{1 / 2}(\\kappa)},\n$$\nwhere \n\n\n$$\nI_{1 / 2}(x)=\\sqrt{\\frac{2}{\\pi x}} \\sinh (x).\n$$\n\nFinally, $p(\\boldsymbol{x} ; \\boldsymbol{\\mu}, \\kappa)$ reads\n\n$$\np(\\boldsymbol{x} ; \\boldsymbol{\\mu}, \\kappa) = \\frac{\\kappa} { (2\\pi)^{3/2} \\sqrt{2} \\sinh(\\kappa) } \\exp \\left(\\kappa \\boldsymbol{\\mu}^{\\mathrm{T}} \\boldsymbol{x}\\right)\n$$\n\nIn the paper, the loss is actually defined as\n\n$$\n\\mathcal{L}= \\mathbb{E}\\left(- {\\rm log}(p)\\right),\n$$\nso this is how I implement it. ","metadata":{}},{"cell_type":"code","source":"import torch \n\ndef VonMisesFisher3DLoss(outputs, targets, kappa):\n    \"\"\"\n    Computes the Von Mises Fisher 3D loss function.\n\n    Args:\n        outputs: A tensor of size (batch_size, 3) representing the predicted unit vectors.\n        targets: A tensor of size (batch_size, 3) representing the true unit vectors.\n        kappa: A float representing the concentration parameter of the distribution.\n    Returns:\n        The loss value as a tensor of size (1,).\n    \"\"\"\n    # Compute the dot product between outputs and targets\n    dot_product = torch.sum(targets * outputs, dim=1)\n    # Compute the Von Mises Fisher loss\n    loss = (kappa / torch.sinh(kappa)) * torch.exp(kappa * dot_product)\n    loss = -torch.log(loss)\n\n    return torch.mean(loss)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-03-10T09:50:48.443218Z","iopub.execute_input":"2023-03-10T09:50:48.443647Z","iopub.status.idle":"2023-03-10T09:50:48.450775Z","shell.execute_reply.started":"2023-03-10T09:50:48.443600Z","shell.execute_reply":"2023-03-10T09:50:48.449241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# example\nx = torch.tensor([[1, 0, 0], [1, 0, 0], [1, 0, 0]], dtype=torch.float32)\nx_target = torch.tensor([[0.9, 0.43588989, 0], [0.9, 0.43588989, 0], [0.9, 0.43588989, 0]], dtype=torch.float32)\nkappa = torch.tensor([0.99, 0.98, 0.1], dtype=torch.float32)\nVonMisesFisher3DLoss(x, x_target, kappa)","metadata":{"execution":{"iopub.status.busy":"2023-03-10T09:50:48.850793Z","iopub.execute_input":"2023-03-10T09:50:48.851188Z","iopub.status.idle":"2023-03-10T09:50:48.864137Z","shell.execute_reply.started":"2023-03-10T09:50:48.851151Z","shell.execute_reply":"2023-03-10T09:50:48.862649Z"},"trusted":true},"execution_count":null,"outputs":[]}]}