Skip to content

Commit 5891642

Browse files
authored
[train][doc] Fix doc code (#39522) (#39558)
Signed-off-by: Justin Yu <justinvyu@anyscale.com>
1 parent 0cf8793 commit 5891642

1 file changed

Lines changed: 11 additions & 41 deletions

File tree

doc/source/train/doc_code/key_concepts.py

Lines changed: 11 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -1,56 +1,26 @@
11
# flake8: noqa
22
# isort: skip_file
33

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
216

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
269
from ray.train.data_parallel_trainer import DataParallelTrainer
2710

2811

2912
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+
)
4519

4620

4721
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)
5123
)
52-
trainer.fit()
53-
# __session_data_info_end__
5424

5525

5626
# __run_config_start__

0 commit comments

Comments
 (0)