Skip to content

Fix silently wrong targets on MPS with >= 65,536 training rows - #1358

Open
rminla wants to merge 2 commits into
PriorLabs:mainfrom
rminla:fix-mps-target-padding
Open

rminla wants to merge 2 commits into
PriorLabs:mainfrom
rminla:fix-mps-target-padding

Conversation

@rminla

@rminla rminla commented Oct 2, 2026 •

Copy link
Copy Markdown

Issue

Fixes #1357

Motivation and Context

On MPS with torch <= 2.14.1, F.pad(..., value=float("nan")) overwrites its input once the tensor reaches 2^16 rows, even with a zero pad width. _prepare_targets uses it on the train targets, so with >= 65,536 training rows every target came out NaN or another row's value, and predictions collapsed to near-constant with no error (details and repro in the issue).

This builds the NaN test rows with new_full and appends them with torch.cat, which gives the same result on every device. _prepare_targets is duplicated in each single-file architecture, so the change is in all five (v2, v2.5, v2.6, v3, v3.5). The comment can go once torch >= 2.15 is the minimum (fixed in nightly 2.15.0.dev20261002).


Public API Changes

  • No Public API changes
  • Yes, Public API changes (Details below)

How Has This Been Tested?

  • New tests/test_architectures/test_prepare_targets.py: all five architectures, 65,535 / 65,536 / 70,000 train rows, with and without test rows, on CPU and (when available) MPS. On torch 2.14.1 / M4 Max: 25 MPS cases fail before the change and 60/60 pass after. CI has no MPS runner, so the MPS cases only run on Apple Silicon.
  • End to end with TabPFNRegressor (v3.5 weights, device="mps", 1,000 held-out rows): MAE at 65,536 train rows 16.58 -> 7.57, at 70,000 rows 13.70 -> 7.55.

Checklist

  • The changes have been tested locally.
  • Documentation has been updated (if the public API or usage changes). Not applicable.
  • A changelog entry has been added (see changelog/README.md), or "no changelog needed" label requested.
  • The code follows the project's style guidelines. ruff==0.15.12 format and check show no new findings.
  • I have considered the impact of these changes on the public API.

@CLAassistant

CLAassistant commented Oct 2, 2026 •

Copy link
Copy Markdown

CLA assistant check
All committers have signed the CLA.

@github-actions
github-actions Bot requested a review from alanprior October 2, 2026 20:10
@alanprior

Copy link
Copy Markdown
Contributor

Thank you for the contribution @rminla, it might take us a few more days to properly review it. I'll follow up on this!

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Silently wrong predictions on MPS with >= 65,536 training rows (NaN F.pad in _prepare_targets)

3 participants