Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -92,8 +92,11 @@ class SklearnPredictionOpDesc extends PythonOperatorDescriptor with StandaloneCo
var resultType = AttributeType.STRING
val inputSchema = inputSchemas(operatorInfo.inputPorts(1).id)
if (groundTruthAttribute != "") {
resultType =
inputSchema.attributes.find(attr => attr.getName == groundTruthAttribute).get.getType
if (!inputSchema.containsAttribute(groundTruthAttribute))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Please keep the exact-case lookup here. With input column label and configured name Label, schema validation now succeeds because containsAttribute ignores case. The generated script then fails at in2df.drop("Label", axis=1) with KeyError. I reproduced this with the generated code; label passes as a control. The native Python filter also compares names exactly.

throw new RuntimeException(
s"Ground Truth column '$groundTruthAttribute' is not in the input table"
)
resultType = inputSchema.getAttribute(groundTruthAttribute).getType
}
Map(
operatorInfo.outputPorts.head.id -> inputSchema
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -71,14 +71,15 @@ class SklearnPredictionOpDescSpec extends AnyFlatSpec with Matchers {
.getType shouldBe AttributeType.INTEGER
}

it should "throw when the configured ground-truth attribute is absent from the input schema" in {
it should "name the configured ground-truth attribute when it is absent from the input schema" in {
val d = new SklearnPredictionOpDesc
d.resultAttribute = "prediction"
d.groundTruthAttribute = "missing"
val data = Schema().add("feature", AttributeType.STRING)
intercept[NoSuchElementException] {
val e = intercept[RuntimeException] {
d.getOutputSchemas(Map(PortIdentity(1) -> data))
}
e.getMessage shouldBe "Ground Truth column 'missing' is not in the input table"
}

"SklearnPredictionOpDesc.generatePythonCode" should "emit the model-applying tuple operator" in {
Expand Down
Loading