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

Method test_flag_utils

axlearn/cloud/common/utils_test.py:440–526  ·  view source on GitHub ↗

Tests define_flags and from_flags.

(self)

Source from the content-addressed store, hash-verified

438 self.assertEqual(fv.shared_override, "child-default")
439
440 def test_flag_utils(self):
441 """Tests define_flags and from_flags."""
442
443 class Inner(utils.FlagConfigurable):
444 """An inner config."""
445
446 @config_class
447 class Config(utils.FlagConfigurable.Config):
448 common_value: Required[str] = REQUIRED
449 inner_value: Required[str] = REQUIRED
450
451 @classmethod
452 def define_flags(cls, fv):
453 super().define_flags(fv)
454 common_kwargs = dict(flag_values=fv, allow_override=True)
455 flags.DEFINE_string("common_value", None, "", **common_kwargs)
456 flags.DEFINE_string("inner_value", None, "", **common_kwargs)
457
458 @classmethod
459 def set_defaults(cls, fv):
460 super().set_defaults(fv)
461 fv.set_default("common_value", "child-default")
462
463 # Test that it traverses non FlagConfigurables.
464 class RegularConfigurable(Configurable):
465 """A dummy container config to test traversal."""
466
467 @config_class
468 class Config(Configurable.Config):
469 # Test that it traverses other non-config containers.
470 inner: list[Inner.Config] = [Inner.default_config()]
471
472 class Outer(utils.FlagConfigurable):
473 """An outer config."""
474
475 @config_class
476 class Config(utils.FlagConfigurable.Config):
477 common_value: Required[str] = REQUIRED
478 outer_value: Required[str] = REQUIRED
479 inner_enabled: Optional[bool] = None
480 inner: Optional[Configurable.Config] = RegularConfigurable.default_config()
481
482 @classmethod
483 def define_flags(cls, fv):
484 super().define_flags(fv)
485 common_kwargs = dict(flag_values=fv, allow_override=True)
486 flags.DEFINE_string("common_value", None, "", **common_kwargs)
487 flags.DEFINE_string("outer_value", None, "", **common_kwargs)
488 flags.DEFINE_bool("inner_enabled", None, "", **common_kwargs)
489
490 @classmethod
491 def set_defaults(cls, fv):
492 super().set_defaults(fv)
493 fv.set_default("common_value", "parent-default")
494
495 @classmethod
496 def from_flags(cls, fv, **kwargs):
497 cfg = super().from_flags(fv, **kwargs)

Callers

nothing calls this directly

Calls 4

cloneMethod · 0.80
default_configMethod · 0.45
define_flagsMethod · 0.45
from_flagsMethod · 0.45

Tested by

no test coverage detected