Skip to content

Commit c93c3ce

Browse files
committed
fix dump sl
1 parent e78b40d commit c93c3ce

6 files changed

Lines changed: 18 additions & 12 deletions

File tree

run_train.sh

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,6 @@ CONFIG=${CONFIG:-"llama3_debugmodel"}
3838
COMM_MODE=${COMM_MODE:-""}
3939
NNODES=${NNODES:-${SLURM_JOB_NUM_NODES:-1}}
4040
NODE_RANK=${NODE_RANK:-${SLURM_NODEID:-0}}
41-
TORCHTITAN_DUMP_FOLDER=${TORCHTITAN_DUMP_FOLDER:-"/tmp/torchtitan_train/"}
4241

4342
# xx injects these
4443
TORCHTITAN_ARGS=()
@@ -88,7 +87,7 @@ export REPORTERV2_TRAINING_ID
8887
if [ -n "$COMM_MODE" ]; then
8988
# Communication mode specified: validate configuration or run in debug mode
9089
echo "Running with comm_mode=${COMM_MODE}"
91-
NGPU="${NGPU}" LOCAL_RANK=0 python3 -m torchtitan.train --module ${MODULE} --config ${CONFIG} --dump-folder "${TORCHTITAN_DUMP_FOLDER}" "$@" --comm.mode=${COMM_MODE} --training.steps 1
90+
NGPU="${NGPU}" LOCAL_RANK=0 python3 -m torchtitan.train --module ${MODULE} --config ${CONFIG} "$@" --comm.mode=${COMM_MODE} --training.steps 1
9291
else
9392
if [[ -n "${MASTER_ADDR:-}" ]]; then
9493
RDZV_ENDPOINT="${MASTER_ADDR}:${MASTER_PORT}"
@@ -103,5 +102,5 @@ else
103102
--rdzv_id=${RDZV_ID:-${SLURM_JOB_ID:-$(generate_uuid)}} --rdzv_backend c10d \
104103
--rdzv_endpoint="${RDZV_ENDPOINT}" \
105104
--local-ranks-filter ${LOG_RANK} --role rank --tee 3 \
106-
-m torchtitan.train --module ${MODULE} --config ${CONFIG} --dump-folder "${TORCHTITAN_DUMP_FOLDER}" "$@"
105+
-m torchtitan.train --module ${MODULE} --config ${CONFIG} "$@"
107106
fi

tests/unit_tests/test_checkpoint.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -386,10 +386,11 @@ def test_dcp_save_load_memory_fs_checkpoint(self):
386386
states=self.states,
387387
config=cfg,
388388
sd_adapter=None,
389-
base_folder="",
389+
base_folder="./outputs",
390390
)
391391

392392
try:
393+
self.assertEqual(manager.folder, root)
393394
checkpoint_id = manager._create_checkpoint_id(1)
394395
expected = torch.tensor([3.0])
395396
manager.dcp_save(

torchtitan/components/checkpoint.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
from typing import Any, cast, Literal
1717
from urllib.parse import urlparse
1818

19+
from fsspec.core import split_protocol
1920
import torch
2021
import torch.distributed as dist
2122
import torch.distributed.checkpoint as dcp
@@ -458,7 +459,11 @@ def __init__(
458459
if not self.enable:
459460
return
460461

461-
self.folder = fs.join_path(base_folder, config.folder)
462+
self.folder = (
463+
config.folder
464+
if split_protocol(config.folder)[0] is not None
465+
else fs.join_path(base_folder, config.folder)
466+
)
462467
self.checkpoint_id_format = config.checkpoint_id_format
463468
self.interval = config.interval
464469

torchtitan/experiments/worldmodel/config_registry.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -63,11 +63,10 @@ def worldmodel() -> WorldModelTrainer.Config:
6363
optimizer = default_adamw(lr=2e-4, weight_decay=1e-2)
6464
optimizer.implementation = "fused_opt_states_bf16"
6565
local_world_size, world_size, num_nodes = _world_sizes()
66-
checkpoint_base_folder = _reporterv2_checkpoint_base_folder()
66+
checkpoint_folder = _reporterv2_checkpoint_folder()
6767

6868
return WorldModelTrainer.Config(
6969
hf_assets_path=".",
70-
dump_folder=checkpoint_base_folder or "./outputs",
7170
loss=WorldModelLoss.Config(plan_loss_weight=0.1),
7271
tokenizer=WorldModelTokenizer.Config(
7372
compressor_model=COMPRESSOR_MODEL,
@@ -111,7 +110,7 @@ def worldmodel() -> WorldModelTrainer.Config:
111110
),
112111
checkpoint=WorldModelTorchPackageCheckpointManager.Config(
113112
enable=True,
114-
folder=os.getenv("REPORTERV2_TRAINING_ID") or "checkpoint",
113+
folder=checkpoint_folder,
115114
interval=validation_freq * 5,
116115
async_mode="async",
117116
keep_latest_k=0,
@@ -155,9 +154,10 @@ def _world_sizes() -> tuple[int, int, int]:
155154
return local_world_size, world_size, num_nodes
156155

157156

158-
def _reporterv2_checkpoint_base_folder() -> str:
157+
def _reporterv2_checkpoint_folder() -> str:
159158
host = os.getenv("REPORTERV2_HOST")
160-
return f"{host.rstrip('/')}/checkpoint" if host else ""
159+
folder = os.getenv("REPORTERV2_TRAINING_ID") or "checkpoint"
160+
return f"{host.rstrip('/')}/checkpoint/{folder}" if host else folder
161161

162162

163163
def main() -> None:

torchtitan/experiments/worldmodel/torchpackage_checkpoint.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,8 @@
3636
PACKAGE_NAME = "model.torchpackage"
3737
MODEL_CONFIG_FILE = "_torchpackage_model_config.pt"
3838
STRUCTURED_LOG_DIR = os.getenv(
39-
"TORCHTITAN_STRUCTURED_LOG_DIR", "./outputs/worldmodel_torchpackage_checkpoint"
39+
"TORCHTITAN_STRUCTURED_LOG_DIR",
40+
"/tmp/torchtitan_train/worldmodel_torchpackage_checkpoint",
4041
)
4142
WORLD_MODEL_TORCH_PACKAGE_RECIPE = (
4243
"torchtitan.experiments.worldmodel.torchpackage_checkpoint:"

torchtitan/trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ class Config(Configurable.Config):
7373
(fqn to file mapping), the config.json file, generation_config.json, and tokenizer files.
7474
"""
7575

76-
dump_folder: str = "./outputs"
76+
dump_folder: str = "/tmp/torchtitan_train"
7777
"""Folder to dump job outputs"""
7878

7979
codedir: str = ""

0 commit comments

Comments
 (0)