|
3 | 3 |
|
4 | 4 | # DeepSpeed Team |
5 | 5 |
|
| 6 | +import sys |
| 7 | + |
6 | 8 | from deepspeed.monitor.tensorboard import TensorBoardMonitor |
7 | 9 | from deepspeed.monitor.wandb import WandbMonitor |
8 | 10 | from deepspeed.monitor.csv_monitor import csvMonitor |
9 | 11 | from deepspeed.monitor.config import DeepSpeedMonitorConfig |
10 | 12 | from deepspeed.monitor.comet import CometMonitor |
| 13 | +from deepspeed.monitor.trackio import TrackioMonitor |
| 14 | +from deepspeed.monitor.monitor import MonitorMaster |
11 | 15 |
|
12 | 16 | from unit.common import DistributedTest |
13 | | -from unittest.mock import Mock, patch |
| 17 | +from unittest.mock import Mock, MagicMock, patch |
14 | 18 | from deepspeed.runtime.config import DeepSpeedConfig |
15 | 19 |
|
16 | 20 | import deepspeed.comm as dist |
@@ -164,3 +168,83 @@ def test_empty_comet(self): |
164 | 168 | assert comet_monitor.enabled == defaults.enabled |
165 | 169 | assert comet_monitor.samples_log_interval == defaults.samples_log_interval |
166 | 170 | mock_start.assert_not_called() |
| 171 | + |
| 172 | + |
| 173 | +class TestTrackio(DistributedTest): |
| 174 | + world_size = 2 |
| 175 | + |
| 176 | + def test_trackio(self): |
| 177 | + # trackio is an optional dependency, so we stub the module rather |
| 178 | + # than requiring it to be installed for CI. |
| 179 | + mock_trackio = MagicMock() |
| 180 | + |
| 181 | + config_dict = {"train_batch_size": 2, "trackio": {"enabled": True, "project": "my_project"}} |
| 182 | + ds_config = DeepSpeedConfig(config_dict) |
| 183 | + |
| 184 | + with patch.dict(sys.modules, {"trackio": mock_trackio}): |
| 185 | + trackio_monitor = TrackioMonitor(ds_config.monitor_config.trackio) |
| 186 | + |
| 187 | + assert trackio_monitor.enabled == True |
| 188 | + assert trackio_monitor.project == "my_project" |
| 189 | + |
| 190 | + # trackio.init should only be called on rank 0 |
| 191 | + if dist.get_rank() == 0: |
| 192 | + mock_trackio.init.assert_called_once_with(project="my_project") |
| 193 | + else: |
| 194 | + mock_trackio.init.assert_not_called() |
| 195 | + |
| 196 | + def test_empty_trackio(self): |
| 197 | + mock_trackio = MagicMock() |
| 198 | + |
| 199 | + config_dict = {"train_batch_size": 2, "trackio": {}} |
| 200 | + ds_config = DeepSpeedConfig(config_dict) |
| 201 | + |
| 202 | + with patch.dict(sys.modules, {"trackio": mock_trackio}): |
| 203 | + trackio_monitor = TrackioMonitor(ds_config.monitor_config.trackio) |
| 204 | + |
| 205 | + defaults = DeepSpeedMonitorConfig().trackio |
| 206 | + assert trackio_monitor.enabled == defaults.enabled |
| 207 | + assert trackio_monitor.project == defaults.project |
| 208 | + |
| 209 | + def test_trackio_write_events(self): |
| 210 | + # Verifies write_events() correctly converts 3-tuples into |
| 211 | + # trackio.log() calls with the right step value. |
| 212 | + mock_trackio = MagicMock() |
| 213 | + |
| 214 | + config_dict = {"train_batch_size": 2, "trackio": {"enabled": True, "project": "my_project"}} |
| 215 | + ds_config = DeepSpeedConfig(config_dict) |
| 216 | + |
| 217 | + with patch.dict(sys.modules, {"trackio": mock_trackio}): |
| 218 | + trackio_monitor = TrackioMonitor(ds_config.monitor_config.trackio) |
| 219 | + events = [("Train/Loss", 0.5, 100)] |
| 220 | + trackio_monitor.write_events(events) |
| 221 | + |
| 222 | + if dist.get_rank() == 0: |
| 223 | + mock_trackio.log.assert_called_once_with({"Train/Loss": 0.5}, step=100) |
| 224 | + else: |
| 225 | + mock_trackio.log.assert_not_called() |
| 226 | + |
| 227 | + |
| 228 | +class TestMonitorMasterTrackioWiring(DistributedTest): |
| 229 | + world_size = 2 |
| 230 | + |
| 231 | + def test_trackio_enabled_creates_monitor(self): |
| 232 | + mock_trackio = MagicMock() |
| 233 | + |
| 234 | + config_dict = {"train_batch_size": 2, "trackio": {"enabled": True, "project": "my_project"}} |
| 235 | + ds_config = DeepSpeedConfig(config_dict) |
| 236 | + |
| 237 | + with patch.dict(sys.modules, {"trackio": mock_trackio}): |
| 238 | + monitor_master = MonitorMaster(ds_config.monitor_config) |
| 239 | + |
| 240 | + if dist.get_rank() == 0: |
| 241 | + assert monitor_master.trackio_monitor is not None |
| 242 | + assert isinstance(monitor_master.trackio_monitor, TrackioMonitor) |
| 243 | + else: |
| 244 | + assert monitor_master.trackio_monitor is None |
| 245 | + |
| 246 | + def test_trackio_disabled_skips_monitor(self): |
| 247 | + config_dict = {"train_batch_size": 2, "trackio": {"enabled": False}} |
| 248 | + ds_config = DeepSpeedConfig(config_dict) |
| 249 | + monitor_master = MonitorMaster(ds_config.monitor_config) |
| 250 | + assert monitor_master.trackio_monitor is None |
0 commit comments