From a38261407180ab63b6ca8b51ab6b839197789c2d Mon Sep 17 00:00:00 2001 From: May <79735133+YuCeong-May@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:04:26 +0800 Subject: [PATCH 1/3] fix(train): initialize DeepSpeed mode from top-level config --- funasr/bin/train_ds.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/funasr/bin/train_ds.py b/funasr/bin/train_ds.py index b32881d99..a61428388 100644 --- a/funasr/bin/train_ds.py +++ b/funasr/bin/train_ds.py @@ -83,9 +83,12 @@ def main(**kwargs): if local_rank == 0: tables.print() - use_ddp = world_size > 1 use_fsdp = kwargs.get("use_fsdp", False) use_deepspeed = kwargs.get("use_deepspeed", False) + deepspeed_config = kwargs.get("deepspeed_config", "") + if use_deepspeed and use_fsdp: + raise ValueError("use_deepspeed and use_fsdp cannot be enabled at the same time") + use_ddp = world_size > 1 and not use_deepspeed and not use_fsdp if use_deepspeed: logging.info(f"use_deepspeed: {use_deepspeed}") deepspeed.init_distributed(dist_backend=kwargs.get("backend", "nccl")) @@ -143,6 +146,8 @@ def main(**kwargs): world_size=world_size, use_ddp=use_ddp, use_fsdp=use_fsdp, + use_deepspeed=use_deepspeed, + deepspeed_config=deepspeed_config, device=kwargs["device"], excludes=kwargs.get("excludes", None), output_dir=kwargs.get("output_dir", "./exp"), From 99cd0a3d128f0ec24db23b6fe6f9a6629c49749a Mon Sep 17 00:00:00 2001 From: May <79735133+YuCeong-May@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:16:11 +0800 Subject: [PATCH 2/3] fix(train): normalize distributed trainer configuration --- funasr/bin/train_ds.py | 46 +++++++++++++++++++++++++++++++++++------- 1 file changed, 39 insertions(+), 7 deletions(-) diff --git a/funasr/bin/train_ds.py b/funasr/bin/train_ds.py index a61428388..d1f25973e 100644 --- a/funasr/bin/train_ds.py +++ b/funasr/bin/train_ds.py @@ -40,6 +40,37 @@ except: deepspeed = None +_DISTRIBUTED_TRAIN_CONF_KEYS = ( + "use_ddp", + "use_fsdp", + "use_deepspeed", + "deepspeed_config", +) + + +def _resolve_distributed_config(kwargs, world_size): + """Resolve distributed settings with top-level values taking precedence.""" + train_conf = dict(kwargs.get("train_conf") or {}) + + def get_setting(name, default): + if name in kwargs: + return kwargs[name] + return train_conf.get(name, default) + + use_fsdp = get_setting("use_fsdp", False) + use_deepspeed = get_setting("use_deepspeed", False) + deepspeed_config = get_setting("deepspeed_config", "") + if use_deepspeed and use_fsdp: + raise ValueError("use_deepspeed and use_fsdp cannot be enabled at the same time") + + trainer_conf = { + key: value + for key, value in train_conf.items() + if key not in _DISTRIBUTED_TRAIN_CONF_KEYS + } + use_ddp = world_size > 1 and not use_deepspeed and not use_fsdp + return use_ddp, use_fsdp, use_deepspeed, deepspeed_config, trainer_conf + @hydra.main(config_name=None, version_base=None) def main_hydra(kwargs: DictConfig): @@ -83,12 +114,13 @@ def main(**kwargs): if local_rank == 0: tables.print() - use_fsdp = kwargs.get("use_fsdp", False) - use_deepspeed = kwargs.get("use_deepspeed", False) - deepspeed_config = kwargs.get("deepspeed_config", "") - if use_deepspeed and use_fsdp: - raise ValueError("use_deepspeed and use_fsdp cannot be enabled at the same time") - use_ddp = world_size > 1 and not use_deepspeed and not use_fsdp + ( + use_ddp, + use_fsdp, + use_deepspeed, + deepspeed_config, + trainer_conf, + ) = _resolve_distributed_config(kwargs, world_size) if use_deepspeed: logging.info(f"use_deepspeed: {use_deepspeed}") deepspeed.init_distributed(dist_backend=kwargs.get("backend", "nccl")) @@ -151,7 +183,7 @@ def main(**kwargs): device=kwargs["device"], excludes=kwargs.get("excludes", None), output_dir=kwargs.get("output_dir", "./exp"), - **kwargs.get("train_conf"), + **trainer_conf, ) model = trainer.warp_model(model, **kwargs) From 241dea6352912336cf57d8d0be3f93533c19eb2e Mon Sep 17 00:00:00 2001 From: May <79735133+YuCeong-May@users.noreply.github.com> Date: Wed, 5 Aug 2026 11:17:10 +0800 Subject: [PATCH 3/3] test(train): cover distributed trainer configuration --- tests/test_train_ds_distributed_config.py | 62 +++++++++++++++++++++++ 1 file changed, 62 insertions(+) create mode 100644 tests/test_train_ds_distributed_config.py diff --git a/tests/test_train_ds_distributed_config.py b/tests/test_train_ds_distributed_config.py new file mode 100644 index 000000000..20537d160 --- /dev/null +++ b/tests/test_train_ds_distributed_config.py @@ -0,0 +1,62 @@ +import unittest + +from funasr.bin.train_ds import _resolve_distributed_config + + +class TestDistributedConfig(unittest.TestCase): + def test_nested_train_conf_is_supported(self): + use_ddp, use_fsdp, use_deepspeed, config, trainer_conf = ( + _resolve_distributed_config( + { + "train_conf": { + "use_deepspeed": True, + "deepspeed_config": "nested_ds.json", + "log_interval": 10, + } + }, + world_size=2, + ) + ) + + self.assertFalse(use_ddp) + self.assertFalse(use_fsdp) + self.assertTrue(use_deepspeed) + self.assertEqual(config, "nested_ds.json") + self.assertEqual(trainer_conf, {"log_interval": 10}) + + def test_top_level_values_override_train_conf(self): + use_ddp, use_fsdp, use_deepspeed, config, trainer_conf = ( + _resolve_distributed_config( + { + "use_deepspeed": False, + "deepspeed_config": "top_level_ds.json", + "train_conf": { + "use_deepspeed": True, + "deepspeed_config": "nested_ds.json", + "use_ddp": True, + "log_interval": 10, + }, + }, + world_size=2, + ) + ) + + self.assertTrue(use_ddp) + self.assertFalse(use_fsdp) + self.assertFalse(use_deepspeed) + self.assertEqual(config, "top_level_ds.json") + self.assertEqual(trainer_conf, {"log_interval": 10}) + + def test_deepspeed_and_fsdp_are_mutually_exclusive(self): + with self.assertRaisesRegex(ValueError, "cannot be enabled"): + _resolve_distributed_config( + { + "use_deepspeed": True, + "train_conf": {"use_fsdp": True}, + }, + world_size=2, + ) + + +if __name__ == "__main__": + unittest.main()