(
cfg: SpmdTrainer.Config, *, kv_separator=": ", field_separator="\n"
)
| 80 | |
| 81 | |
| 82 | def param_init_debug_string( |
| 83 | cfg: SpmdTrainer.Config, *, kv_separator=": ", field_separator="\n" |
| 84 | ) -> str: |
| 85 | layer = cfg.model.set(name="init_debug_string").instantiate(parent=None) |
| 86 | param_init_specs = read_param_init_specs_recursively( |
| 87 | layer, |
| 88 | delegates={ |
| 89 | HF_MODULE_KEY: ParamInitSpec( |
| 90 | shape=None, |
| 91 | initializer=ThirdPartyInitializer.default_config().set(library="hf"), |
| 92 | fan_axes=None, |
| 93 | ), |
| 94 | }, |
| 95 | ) |
| 96 | lines = [] |
| 97 | for name, init_specs in flatten_items(param_init_specs): |
| 98 | init_str = init_specs.initializer.debug_string( |
| 99 | name=name, shape=init_specs.shape, axes=init_specs.fan_axes |
| 100 | ) |
| 101 | lines.append(f"{name}{kv_separator}{init_str}") |
| 102 | return field_separator.join(lines) |
| 103 | |
| 104 | |
| 105 | def per_param_setting_debug_string( |
no test coverage detected