Skip to content

Issue 2808 - Fix sklearn .fit method incorrectly applies arg max bug - #2811

Draft
wasibabi wants to merge 1 commit into
Trusted-AI:mainfrom
wasibabi:issue/2808/skfit_learn_method/arg_max_bug
Draft

wasibabi wants to merge 1 commit into
Trusted-AI:mainfrom
wasibabi:issue/2808/skfit_learn_method/arg_max_bug

Conversation

@wasibabi

@wasibabi wasibabi commented Jun 12, 2026 •

Copy link
Copy Markdown

Description

This is a proposed fix to issue 2808

  • The issue is that ScikitlearnClassifier.fit unconditionally applies np.argmax. This works when y are one-hot encoded labels but fails for 1D vectors of class indices. Based on the docs I believe ScikitlearnClassifier.fit is supposed to support both
  • The fix is to conditionally apply np.argmax only for one-hot encoded lables

Type of change

Please check all relevant options.

  • Improvement (non-breaking)
  • Bug fix (non-breaking)
  • New feature (non-breaking)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • This change requires a documentation update

Testing

Please describe the tests that you ran to verify your changes. Consider listing any relevant details of your test configuration.

  • Reproduce the bug on main
    • with y = np.array([0, 0, 0, 1, 1, 1])
    • run clf.fit(...)
    • Get `numpy.exceptions.AxisError: axis 1 is
art issue 2808 - bug repro - class indices error out of bounds for array of dimension 1`
  • Reproduce one-hot encoding working on main
    • with y = np.array([[1, 0], [1, 0], [1, 0], [0, 1], [0, 1], [0, 1]])
    • run clf.fit(...)
    • run succeeds. Can print out clf.nb_classes and clf.model.classes
art issue 2808 - bug repro - one hot encoding
  • Demo fix on bugfix branch
    • same test
art issue 2808 - bug fix demo - class indices
  • Demo bugfix branch handles one-hot encoding
    • same test
art issue 2808 - bug fix demo - one hot encoding

Test Configuration:

  • OS: MacOS Tahoe
  • Python version: 3.11.8
  • ART version or commit number: latest master
  • TensorFlow / Keras / PyTorch version

Checklist

  • My code follows the style guidelines of this project
  • I have performed a self-review of my own code
  • I have commented my code
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • My changes have been tested using both CPU and GPU devices

Signed-off-by: Dylan Jones <jonesdylan038@gmail.com>
@wasibabi
wasibabi force-pushed the issue/2808/skfit_learn_method/arg_max_bug branch from 6f704e5 to da5f70f Compare June 12, 2026 20:08
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.

1 participant