Skip to content

Commit 1501053

Browse files
Merge pull request #855 from DashAISoftware/fix/regression-prediction-crashes
Fix regression prediction crashes
2 parents 7c1a0d6 + 3953cb9 commit 1501053

2 files changed

Lines changed: 15 additions & 0 deletions

File tree

DashAI/back/job/predict_job.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
from DashAI.back.job.base_job import BaseJob, JobError
1313
from DashAI.back.models.base_model import BaseModel
1414
from DashAI.back.tasks.base_task import BaseTask
15+
from DashAI.back.tasks.regression_task import RegressionTask
1516

1617
if TYPE_CHECKING:
1718
from sqlalchemy.orm import sessionmaker
@@ -482,6 +483,14 @@ def run(
482483
if key in model_session.input_columns + model_session.output_columns
483484
}
484485

486+
# Regression models predict continuous values regardless of
487+
# the training target's original dtype (e.g. a target column
488+
# that happened to hold only integer-looking values), so the
489+
# output column's saved schema must reflect that instead of
490+
# inheriting the training dataset's type.
491+
if isinstance(task, RegressionTask):
492+
filtered_schema[output_col] = {"type": "Float", "dtype": "float64"}
493+
485494
# Store num of rows, columns, and column names
486495
dataset_with_prediction.compute_base_metadata()
487496

DashAI/back/models/scikit_learn/sklearn_like_model.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -79,6 +79,12 @@ def train(self, x_train, y_train, x_validation=None, y_validation=None):
7979
"""
8080
x_processed = self.prepare_dataset(x_train, is_fit=True).to_pandas()
8181
y_processed = self.prepare_output(y_train, is_fit=True).to_pandas()
82+
# Every task using this base class has outputs_cardinality 1, so this
83+
# is always a single column. Passed as a DataFrame, some estimators
84+
# (e.g. LinearRegression) keep predictions 2D to match; squeezing to a
85+
# Series here keeps fit/predict shapes 1D for every estimator alike.
86+
if y_processed.shape[1] == 1:
87+
y_processed = y_processed.iloc[:, 0]
8288
return super().fit(x_processed, y_processed)
8389

8490
def predict(self, x: "DashAIDataset"):

0 commit comments

Comments
 (0)