From 1eedfc642e251570703ba5df767a93afc4c9dcef Mon Sep 17 00:00:00 2001 From: Mathieu Guillame-Bert Date: Thu, 1 Oct 2026 06:06:13 -0700 Subject: [PATCH] Return error when training a node classification model on negative integer labels. PiperOrigin-RevId: 991627689 --- .../ten_lines/node_prediction_test.py | 58 +++++++++++++++++++ .../ten_lines/node_prediction_train.py | 17 ++++++ 2 files changed, 75 insertions(+) diff --git a/dgf/src/learning/ten_lines/node_prediction_test.py b/dgf/src/learning/ten_lines/node_prediction_test.py index 0c8ebcb..0916abd 100644 --- a/dgf/src/learning/ten_lines/node_prediction_test.py +++ b/dgf/src/learning/ten_lines/node_prediction_test.py @@ -996,6 +996,64 @@ def test_label_classes_fails_for_integer_labels(self): self.assertIn("does not have a string dictionary", str(ctx.exception)) +class NodePredictionNegativeLabelTest(absltest.TestCase): + """Integer classification labels must be non-negative. + + A negative label (e.g. -1) generally encodes a missing label. + """ + + @classmethod + def setUpClass(cls): + super().setUpClass() + cls.graph, cls.schema = gen_test_graph.gen_toy_classification_dataset() + # The first 10 nodes have a missing label. + cls.graph.node_sets["N1"].features["label"][:10] = -1 + + def test_in_memory_negative_label_raises(self): + with self.assertRaisesRegex( + ValueError, "should be non-negative. However, its minimum value is -1" + ): + node_prediction_lib.train_node_model( + graph=self.graph, + schema=self.schema, + target_nodeset="N1", + target_column="label", + **RAPID_TRAINING_KWARGS, + ) + + def test_pre_sampled_negative_label_raises(self): + sampler = in_memory_sampler_lib.create_sampler( + graph=self.graph, + plan=sampling_config_lib.simple_sampling_config_to_sampling_plan( + sampling_config_lib.SimpleSamplingConfig( + seed_nodeset="N1", num_hops=1, hop_width=3, reverse=True + ), + self.schema, + ), + schema=self.schema, + batch_size=10, + ) + path = os.path.join(self.create_tempdir().full_path, "samples@2.tfrecord") + tf_graph_sample_lib.write_tfgnn_graphs( + (sample for sample in sampler.sample(np.arange(10))), + path, + schema=self.schema, + container_type="TF_RECORD", + ) + + with self.assertRaisesRegex( + ValueError, "should be non-negative. However, its minimum value is -1" + ): + node_prediction_lib.train_node_model( + graph=path, + valid_graph=path, + schema=self.schema, + target_nodeset="N1", + target_column="label", + **RAPID_TRAINING_KWARGS, + ) + + class NodePredictionRegressionToy(parameterized.TestCase): @parameterized.parameters(("float32", 1), ("int64", 1), ("float32", 3)) diff --git a/dgf/src/learning/ten_lines/node_prediction_train.py b/dgf/src/learning/ten_lines/node_prediction_train.py index e495b4d..75ce30c 100644 --- a/dgf/src/learning/ten_lines/node_prediction_train.py +++ b/dgf/src/learning/ten_lines/node_prediction_train.py @@ -466,6 +466,23 @@ def train_node_model( f" '{task.target_nodeset}' must have `num_categorical_values`" " defined for classification tasks." ) + if ( + schema.node_sets[task.target_nodeset] + .features[task.target_column] + .format.is_integer() + ): + label_stats = dataset_preparator.feature_stats.node_sets[ + task.target_nodeset + ].features[task.target_column] + if label_stats.minimum < 0: + raise ValueError( + f"Target column '{task.target_column}' in nodeset" + f" '{task.target_nodeset}' is an integer categorical feature, so" + " its values should be non-negative. However, its minimum value" + f" is {label_stats.minimum:g}. Negative values generally indicate" + " missing labels. Make sure that all the nodes have a valid" + " label." + ) elif task.task_type == node_prediction_model.TaskType.NODE_REGRESSION: if label_spec.semantic != schema_lib.FeatureSemantic.NUMERICAL: raise ValueError(