@@ -1044,6 +1044,75 @@ class Test(ConfigBase):
10441044 except ImportError :
10451045 pass
10461046
1047+ def test_skip_serialization_metadata (self ):
1048+ """Tests that dataclass fields with skip_serialization metadata are excluded
1049+ from serialization.
1050+ """
1051+
1052+ @dataclasses .dataclass
1053+ class DataWithNonTrainingFields :
1054+ """Test class that has fields excluded from serialization"""
1055+
1056+ important_field : str
1057+ # Field with skip_serialization metadata should be excluded
1058+ internal_cache : dict = dataclasses .field (
1059+ default_factory = dict , metadata = {"skip_serialization" : True }
1060+ )
1061+ # Another field with skip_serialization metadata
1062+ debug_info : str = dataclasses .field (
1063+ default = "debug" , metadata = {"skip_serialization" : True }
1064+ )
1065+ # Regular field without metadata
1066+ regular_field : int = 42
1067+
1068+ @config_class
1069+ class TestConfig (ConfigBase ):
1070+ data : DataWithNonTrainingFields = DataWithNonTrainingFields (
1071+ important_field = "test"
1072+ ) # pytype: disable=invalid-annotation
1073+
1074+ cfg = TestConfig ()
1075+
1076+ # Test with omit_default_values=set() to ensure fields are excluded regardless of defaults
1077+ flat_dict = cfg .to_flat_dict (omit_default_values = set ())
1078+
1079+ # Check that skip_serialization fields are excluded
1080+ self .assertNotIn ("data['internal_cache']" , flat_dict )
1081+ self .assertNotIn ("data['debug_info']" , flat_dict )
1082+
1083+ # Check that regular fields are included
1084+ self .assertIn ("data['important_field']" , flat_dict )
1085+ self .assertIn ("data['regular_field']" , flat_dict )
1086+ self .assertEqual (flat_dict ["data['important_field']" ], "test" )
1087+ self .assertEqual (flat_dict ["data['regular_field']" ], 42 )
1088+
1089+ # Test debug_string as well
1090+ debug_str = cfg .debug_string (omit_default_values = set ())
1091+ self .assertNotIn ("internal_cache" , debug_str )
1092+ self .assertNotIn ("debug_info" , debug_str )
1093+ self .assertIn ("important_field" , debug_str )
1094+ self .assertIn ("regular_field" , debug_str )
1095+
1096+ # Test that fields are excluded even when values are assigned
1097+ cfg .data = DataWithNonTrainingFields (
1098+ important_field = "updated" ,
1099+ internal_cache = {"key" : "value" }, # Assign a value
1100+ debug_info = "updated debug" , # Assign a value
1101+ regular_field = 100 ,
1102+ )
1103+
1104+ flat_dict = cfg .to_flat_dict (omit_default_values = set ())
1105+
1106+ # Even with values assigned, skip_serialization fields should still be excluded
1107+ self .assertNotIn ("data['internal_cache']" , flat_dict )
1108+ self .assertNotIn ("data['debug_info']" , flat_dict )
1109+
1110+ # Regular fields should still be included with updated values
1111+ self .assertIn ("data['important_field']" , flat_dict )
1112+ self .assertIn ("data['regular_field']" , flat_dict )
1113+ self .assertEqual (flat_dict ["data['important_field']" ], "updated" )
1114+ self .assertEqual (flat_dict ["data['regular_field']" ], 100 )
1115+
10471116
10481117if __name__ == "__main__" :
10491118 # we can’t use absltest.main because it sets __module__ to "" instead of "axlearn.common", and
0 commit comments