|
1 | 1 | # flake8: noqa |
2 | 2 | # isort: skip_file |
3 | 3 |
|
4 | | -# __session_report_start__ |
5 | | -from ray import train |
6 | | -from ray.train.data_parallel_trainer import DataParallelTrainer |
7 | | - |
8 | | - |
9 | | -def train_fn(config): |
10 | | - for i in range(10): |
11 | | - train.report({"step": i}) |
12 | | - |
13 | | - |
14 | | -trainer = DataParallelTrainer( |
15 | | - train_loop_per_worker=train_fn, scaling_config=train.ScalingConfig(num_workers=1) |
16 | | -) |
17 | | -trainer.fit() |
18 | | - |
19 | | -# __session_report_end__ |
20 | | - |
| 4 | +from pathlib import Path |
| 5 | +import tempfile |
21 | 6 |
|
22 | | -# __session_data_info_start__ |
23 | | -import ray.data |
24 | | - |
25 | | -from ray.train import ScalingConfig |
| 7 | +from ray import train |
| 8 | +from ray.train import Checkpoint |
26 | 9 | from ray.train.data_parallel_trainer import DataParallelTrainer |
27 | 10 |
|
28 | 11 |
|
29 | 12 | def train_fn(config): |
30 | | - context = ray.train.get_context() |
31 | | - dataset_shard = train.get_dataset_shard("train") |
32 | | - |
33 | | - ray.train.report( |
34 | | - { |
35 | | - # Global world size |
36 | | - "world_size": context.get_world_size(), |
37 | | - # Global worker rank on the cluster |
38 | | - "world_rank": context.get_world_rank(), |
39 | | - # Local worker rank on the current machine |
40 | | - "local_rank": context.get_local_rank(), |
41 | | - # Data |
42 | | - "data_shard": next(iter(dataset_shard.iter_batches(batch_format="pandas"))), |
43 | | - } |
44 | | - ) |
| 13 | + for i in range(3): |
| 14 | + with tempfile.TemporaryDirectory() as temp_checkpoint_dir: |
| 15 | + Path(temp_checkpoint_dir).joinpath("model.pt").touch() |
| 16 | + train.report( |
| 17 | + {"loss": i}, checkpoint=Checkpoint.from_directory(temp_checkpoint_dir) |
| 18 | + ) |
45 | 19 |
|
46 | 20 |
|
47 | 21 | trainer = DataParallelTrainer( |
48 | | - train_loop_per_worker=train_fn, |
49 | | - scaling_config=ScalingConfig(num_workers=2), |
50 | | - datasets={"train": ray.data.from_items([1, 2, 3, 4])}, |
| 22 | + train_fn, scaling_config=train.ScalingConfig(num_workers=2) |
51 | 23 | ) |
52 | | -trainer.fit() |
53 | | -# __session_data_info_end__ |
54 | 24 |
|
55 | 25 |
|
56 | 26 | # __run_config_start__ |
|
0 commit comments