Skip to content
Merged
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
58 changes: 58 additions & 0 deletions dgf/src/learning/ten_lines/node_prediction_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
17 changes: 17 additions & 0 deletions dgf/src/learning/ten_lines/node_prediction_train.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading