Resolve model.mode for every output_schema form (fixes Trainer.evaluate crash, #914) - #1246
Merged
Merged
Conversation
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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Problem
A task's
output_schemaentry can be a string ("multilabel"), a processor class (MultiLabelProcessor), a processor instance, or a(name, kwargs)tuple.SampleBuilderaccepts 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:With the string form this happens to work. With the class form,
model.modebecomes a class, which causes three failures:Trainer.evaluate()raisesValueError: Mode <class '...MultiLabelProcessor'> is not supported. This hits built-in tasks such asMIMIC3ICD9Coding, whose schema usesMultiLabelProcessor. Training runs fine and the crash only appears at evaluation.model.mode == "multiclass", which is now False, so they choose the wrong label dtype.TemperatureScalingfails withexpected scalar type Long but found Float.Agent.forwardraisesUnsupported mode.The tuple form fails for every model, because
_resolve_modedoesn't handle tuples.The existing tests with class schemas only exercise
forward().forward()calls_resolve_modeitself, so it passes, and no test ranTrainer.evaluate()with a non-string schema.Fix
BaseModel.modeis now a property. Its setter resolves any assigned value tobinary/multiclass/multilabel/regression. An unresolvable value givesNoneand 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 oldmodel.modeattribute keeps working as before._resolve_modeaccepts the(name or class, kwargs)tuple form.With only the
BaseModelchange, before any model lines were removed, every model in the new test matrix already passed.Tests
tests/core/test_model_mode.py:Trainer.evaluate()with a multiclass label.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
Trainerpass: 279 tests across 29 files. The full core suite also passes:Ran 1332 tests … OK (skipped=76).Docs and example
docs/api/models.rstexplains that the schema forms are equivalent and whatmodel.modecontains.examples/label_schema_forms.pytrains and evaluates an RNN with each form on synthetic data (a few seconds on CPU). Onmasterit crashes on the class form; with this PR, all three forms agree.BaseModeldocstring has a>>>example.Relation to #914
This fixes the crash reported in #914. The broader change suggested there is to have
Trainerand calibration read the label kind from the fitted output processor instead ofmodel.mode. I've proposed it separately on #914 for discussion. It can build on this PR: the newmodeproperty can become a derived alias without changing callers.🤖 Generated with Claude Code