AMP failing on specific line of code
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 9k
- Forks
- 1.5k
- Avg merge
- 2d 4h
- Merged PRs (30d)
- 3
Description
I am using amp with an opt_level="O1". It fails on a specific line of code, saying it expected argument #2 to be of type half but found type float
max_val = torch.max(pos_pairs, torch.max(neg_pairs, dim=1, keepdim=True)[0])
The line comes from here
Any idea why amp is failing to handle this line of code properly? I am using opt_level="01".
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start with the referenced line in pytorch_metric_learning/losses/ntxent_loss.py and reproduce it using AMP with opt_level="O1". Inspect the reported mismatch between half and float arguments to torch.max; the issue is done when the cause is documented and a compatible behavior or workaround is confirmed.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning
- Issue type
- Bug
- Difficulty
- 4/5
- Estimated time
- 3-5 days
- Activity status
- Stale
- Clarity
- Needs clarification
- Newbie friendliness
- 25/100