MCPcopy Create free account
hub / github.com/apple/axlearn / test_set_defaults

Method test_set_defaults

axlearn/cloud/common/utils_test.py:389–438  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

387 """Tests FlagConfigurable."""
388
389 def test_set_defaults(self):
390 class Parent(utils.FlagConfigurable):
391 """A parent class."""
392
393 @classmethod
394 def define_flags(cls, fv):
395 flags.DEFINE_string("shared", None, "", flag_values=fv, allow_override=True)
396 flags.DEFINE_string("parent_only", None, "", flag_values=fv, allow_override=True)
397
398 @classmethod
399 def set_defaults(cls, fv: flags.FlagValues):
400 super().set_defaults(fv)
401 fv.set_default("shared", "parent-default")
402 fv.set_default("parent_only", "parent-default")
403 fv.set_default("shared_override", "parent-default")
404
405 class Child(Parent):
406 """A child class."""
407
408 @classmethod
409 def define_flags(cls, fv):
410 super().define_flags(fv)
411 flags.DEFINE_string("shared", None, "", flag_values=fv, allow_override=True)
412 flags.DEFINE_string("child_only", None, "", flag_values=fv, allow_override=True)
413 flags.DEFINE_string(
414 "shared_override", None, "", flag_values=fv, allow_override=True
415 )
416
417 @classmethod
418 def set_defaults(cls, fv: flags.FlagValues):
419 super().set_defaults(fv)
420 fv.set_default("shared_override", "child-default")
421 fv.set_default("child_only", "child-default")
422
423 fv = flags.FlagValues()
424 cfg = Child.default_config()
425 utils.define_flags(cfg, fv)
426 fv.mark_as_parsed()
427 utils.from_flags(cfg, fv)
428
429 # "parent-only" and "child-only" should follow original defaults.
430 self.assertEqual(fv.parent_only, "parent-default")
431 self.assertEqual(fv.child_only, "child-default")
432
433 # "shared" should follow parent default, because it is not overridden.
434 # Note that this is the case even though child defines the same flag.
435 self.assertEqual(fv.shared, "parent-default")
436
437 # "shared_override" should follow child default, because it is overridden.
438 self.assertEqual(fv.shared_override, "child-default")
439
440 def test_flag_utils(self):
441 """Tests define_flags and from_flags."""

Callers

nothing calls this directly

Calls 3

default_configMethod · 0.45
define_flagsMethod · 0.45
from_flagsMethod · 0.45

Tested by

no test coverage detected