Skip to content

Commit f194898

Browse files
Muzhi Zhaochanglan
authored andcommitted
Add flag to exclude non training fields in training config
GitOrigin-RevId: 6b5ad80
1 parent ae23fdb commit f194898

2 files changed

Lines changed: 77 additions & 0 deletions

File tree

axlearn/common/config.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -511,19 +511,27 @@ def to_flat_dict(self, *, omit_default_values: Collection[Any]) -> dict[str, Any
511511
def enter(key: str, val: Any, default_result: Optional[list]) -> Optional[list]:
512512
if dataclasses.is_dataclass(val) and not isinstance(val, type):
513513
fields_default_dict = {}
514+
cur_key_to_field_dict = {}
514515
for field in dataclasses.fields(val):
515516
# Concatenate field name to key as the full field key name.
516517
# Eg: key="my_config.cats[0]", field.name="adopted"
517518
# cur_key="my_config.cats[0]['adopted']"
518519
cur_key = f"{key}['{field.name}']"
519520
fields_default_dict[cur_key] = field.default
521+
cur_key_to_field_dict[cur_key] = field
520522

521523
kvs_to_traverse = []
522524
for cur_key, cur_val in default_result:
523525
if cur_key not in fields_default_dict:
524526
raise KeyError(
525527
f"Field name {cur_key} is not found for dataclass type value."
526528
)
529+
# Get the field to check metadata
530+
field = cur_key_to_field_dict[cur_key]
531+
# Skip fields marked with skip_serialization=True in metadata
532+
if field and field.metadata.get("skip_serialization", False):
533+
continue
534+
527535
default_val = fields_default_dict[cur_key]
528536
if cur_val is default_val and default_val in omit_default_values:
529537
continue

axlearn/common/config_test.py

Lines changed: 69 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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

10481117
if __name__ == "__main__":
10491118
# we can’t use absltest.main because it sets __module__ to "" instead of "axlearn.common", and

0 commit comments

Comments
 (0)