Skip to content

Commit 112b80f

Browse files
authored
fix(train): initialize DeepSpeed mode from resolved config (#3471)
Resolve distributed trainer settings once with top-level precedence and backwards-compatible train_conf fallbacks. Remove consumed mode keys before expanding trainer_conf, reject simultaneous FSDP and DeepSpeed, and cover the behavior with focused tests.
1 parent d410a56 commit 112b80f

2 files changed

Lines changed: 103 additions & 4 deletions

File tree

funasr/bin/train_ds.py

Lines changed: 41 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,37 @@
4040
except:
4141
deepspeed = None
4242

43+
_DISTRIBUTED_TRAIN_CONF_KEYS = (
44+
"use_ddp",
45+
"use_fsdp",
46+
"use_deepspeed",
47+
"deepspeed_config",
48+
)
49+
50+
51+
def _resolve_distributed_config(kwargs, world_size):
52+
"""Resolve distributed settings with top-level values taking precedence."""
53+
train_conf = dict(kwargs.get("train_conf") or {})
54+
55+
def get_setting(name, default):
56+
if name in kwargs:
57+
return kwargs[name]
58+
return train_conf.get(name, default)
59+
60+
use_fsdp = get_setting("use_fsdp", False)
61+
use_deepspeed = get_setting("use_deepspeed", False)
62+
deepspeed_config = get_setting("deepspeed_config", "")
63+
if use_deepspeed and use_fsdp:
64+
raise ValueError("use_deepspeed and use_fsdp cannot be enabled at the same time")
65+
66+
trainer_conf = {
67+
key: value
68+
for key, value in train_conf.items()
69+
if key not in _DISTRIBUTED_TRAIN_CONF_KEYS
70+
}
71+
use_ddp = world_size > 1 and not use_deepspeed and not use_fsdp
72+
return use_ddp, use_fsdp, use_deepspeed, deepspeed_config, trainer_conf
73+
4374

4475
@hydra.main(config_name=None, version_base=None)
4576
def main_hydra(kwargs: DictConfig):
@@ -83,9 +114,13 @@ def main(**kwargs):
83114
if local_rank == 0:
84115
tables.print()
85116

86-
use_ddp = world_size > 1
87-
use_fsdp = kwargs.get("use_fsdp", False)
88-
use_deepspeed = kwargs.get("use_deepspeed", False)
117+
(
118+
use_ddp,
119+
use_fsdp,
120+
use_deepspeed,
121+
deepspeed_config,
122+
trainer_conf,
123+
) = _resolve_distributed_config(kwargs, world_size)
89124
if use_deepspeed:
90125
logging.info(f"use_deepspeed: {use_deepspeed}")
91126
deepspeed.init_distributed(dist_backend=kwargs.get("backend", "nccl"))
@@ -143,10 +178,12 @@ def main(**kwargs):
143178
world_size=world_size,
144179
use_ddp=use_ddp,
145180
use_fsdp=use_fsdp,
181+
use_deepspeed=use_deepspeed,
182+
deepspeed_config=deepspeed_config,
146183
device=kwargs["device"],
147184
excludes=kwargs.get("excludes", None),
148185
output_dir=kwargs.get("output_dir", "./exp"),
149-
**kwargs.get("train_conf"),
186+
**trainer_conf,
150187
)
151188

152189
model = trainer.warp_model(model, **kwargs)
Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
import unittest
2+
3+
from funasr.bin.train_ds import _resolve_distributed_config
4+
5+
6+
class TestDistributedConfig(unittest.TestCase):
7+
def test_nested_train_conf_is_supported(self):
8+
use_ddp, use_fsdp, use_deepspeed, config, trainer_conf = (
9+
_resolve_distributed_config(
10+
{
11+
"train_conf": {
12+
"use_deepspeed": True,
13+
"deepspeed_config": "nested_ds.json",
14+
"log_interval": 10,
15+
}
16+
},
17+
world_size=2,
18+
)
19+
)
20+
21+
self.assertFalse(use_ddp)
22+
self.assertFalse(use_fsdp)
23+
self.assertTrue(use_deepspeed)
24+
self.assertEqual(config, "nested_ds.json")
25+
self.assertEqual(trainer_conf, {"log_interval": 10})
26+
27+
def test_top_level_values_override_train_conf(self):
28+
use_ddp, use_fsdp, use_deepspeed, config, trainer_conf = (
29+
_resolve_distributed_config(
30+
{
31+
"use_deepspeed": False,
32+
"deepspeed_config": "top_level_ds.json",
33+
"train_conf": {
34+
"use_deepspeed": True,
35+
"deepspeed_config": "nested_ds.json",
36+
"use_ddp": True,
37+
"log_interval": 10,
38+
},
39+
},
40+
world_size=2,
41+
)
42+
)
43+
44+
self.assertTrue(use_ddp)
45+
self.assertFalse(use_fsdp)
46+
self.assertFalse(use_deepspeed)
47+
self.assertEqual(config, "top_level_ds.json")
48+
self.assertEqual(trainer_conf, {"log_interval": 10})
49+
50+
def test_deepspeed_and_fsdp_are_mutually_exclusive(self):
51+
with self.assertRaisesRegex(ValueError, "cannot be enabled"):
52+
_resolve_distributed_config(
53+
{
54+
"use_deepspeed": True,
55+
"train_conf": {"use_fsdp": True},
56+
},
57+
world_size=2,
58+
)
59+
60+
61+
if __name__ == "__main__":
62+
unittest.main()

0 commit comments

Comments
 (0)