Download loss/mae_loss.py from adelelsayed1991/mae: direct link, hf CLI and curl.
- Browser
- Download file 328 Bytes
-
https://huggingface.co/adelelsayed1991/mae/resolve/main/loss/mae_loss.py
- Command line
-
hf download hf://adelelsayed1991/mae/loss/mae_loss.py
-
curl -L -o mae_loss.py https://huggingface.co/adelelsayed1991/mae/resolve/main/loss/mae_loss.py
328 Bytes
| import torch | |
| import torch.nn as nn | |
| def mae_loss(pred, target, mask): | |
| # pred/target: (B, N, P), mask: (B, N) with 1=masked | |
| B, N, P = pred.shape | |
| mask = mask.unsqueeze(-1).float() # (B, N, 1) | |
| loss = (pred - target) ** 2 | |
| loss = (loss * mask).sum() / mask.sum().clamp_min(1.0) | |
| return loss |