Skip to content

Resolve model.mode for every output_schema form (fixes Trainer.evaluate crash, #914) - #1246

Merged
jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:fix/model-mode-overwrite
Sep 27, 2026
Merged

jhnwu3 merged 1 commit into
sunlabuiuc:masterfrom
solarsys:fix/model-mode-overwrite

Conversation

@solarsys

Copy link
Copy Markdown
Collaborator

Problem

A task's output_schema entry can be a string ("multilabel"), a processor class (MultiLabelProcessor), a processor instance, or a (name, kwargs) tuple. SampleBuilder accepts all of them. BaseModel.__init__ resolves the entry to a mode string (added in #564), but 16 models then overwrite it with the raw entry:

self.mode = self.dataset.output_schema[self.label_key]

With the string form this happens to work. With the class form, model.mode becomes a class, which causes three failures:

  • Trainer.evaluate() raises ValueError: Mode <class '...MultiLabelProcessor'> is not supported. This hits built-in tasks such as MIMIC3ICD9Coding, whose schema uses MultiLabelProcessor. Training runs fine and the crash only appears at evaluation.
  • Calibration methods compare model.mode == "multiclass", which is now False, so they choose the wrong label dtype. TemperatureScaling fails with expected scalar type Long but found Float.
  • Agent.forward raises Unsupported mode.

The tuple form fails for every model, because _resolve_mode doesn't handle tuples.

The existing tests with class schemas only exercise forward(). forward() calls _resolve_mode itself, so it passes, and no test ran Trainer.evaluate() with a non-string schema.

Fix

  • BaseModel.mode is now a property. Its setter resolves any assigned value to binary / multiclass / multilabel / regression. An unresolvable value gives None and logs a warning. As a result, a model that assigns the raw entry, including third-party subclasses, still ends up with a string, and the pattern can't come back. The old model.mode attribute keeps working as before.
  • _resolve_mode accepts the (name or class, kwargs) tuple form.
  • The 19 redundant raw overwrites in 16 models are removed. The lines in medfuse, safedrug and molerec already resolve correctly and are unchanged.

With only the BaseModel change, before any model lines were removed, every model in the new test matrix already passed.

Tests

tests/core/test_model_mode.py:

  • 21 models × 3 schema forms (string / class / tuple), each run through Trainer.evaluate() with a multiclass label.
  • MICRON, which is multilabel-only, with the three multilabel forms.
  • The setter: every form resolves, and non-label values warn and give None.
  • TemperatureScaling.calibrate() with a processor-class schema.

Before the fix: 39 failures and 3 errors. After: all pass. The existing tests for all affected models, calibration and Trainer pass: 279 tests across 29 files. The full core suite also passes: Ran 1332 tests … OK (skipped=76).

Docs and example

  • docs/api/models.rst explains that the schema forms are equivalent and what model.mode contains.
  • examples/label_schema_forms.py trains and evaluates an RNN with each form on synthetic data (a few seconds on CPU). On master it crashes on the class form; with this PR, all three forms agree.
  • The BaseModel docstring has a >>> example.

Relation to #914

This fixes the crash reported in #914. The broader change suggested there is to have Trainer and calibration read the label kind from the fitted output processor instead of model.mode. I've proposed it separately on #914 for discussion. It can build on this PR: the new mode property can become a derived alias without changing callers.

🤖 Generated with Claude Code

output_schema entries may be a string, a processor class, a processor
instance or a (name, kwargs) tuple. BaseModel resolved the mode to a string,
but 16 models then overwrote it with the raw schema entry
(`self.mode = self.dataset.output_schema[self.label_key]`). With a processor
class schema (e.g. MIMIC3ICD9Coding's MultiLabelProcessor), Trainer.evaluate
raised "Mode <class ...> is not supported", calibration methods picked the
wrong label dtype, and Agent raised in forward. The tuple form, which
SampleBuilder accepts, failed for every model.

- BaseModel.mode is now a property; its setter resolves any assigned value
  to one of binary/multiclass/multilabel/regression (or None with a
  warning), so subclasses cannot reintroduce the bug.
- _resolve_mode accepts the (name or class, kwargs) tuple form.
- Remove the 19 redundant raw overwrites in 16 models.
- tests: 21 models x 3 schema forms through Trainer.evaluate, MICRON with
  multilabel forms, the setter, and TemperatureScaling with a class schema.
- docs/api/models.rst: document the equivalent schema forms and model.mode.
- examples/label_schema_forms.py: train/evaluate RNN with each form.

Fixes the crash reported in sunlabuiuc#914. The broader redesign proposed there
(reading the label kind from the fitted output processor) is left for a
follow-up.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

@jhnwu3 jhnwu3 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

lgtm

@jhnwu3
jhnwu3 merged commit f6a8a93 into sunlabuiuc:master Sep 27, 2026
2 checks passed
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.

2 participants