Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions run_train.sh
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@ export LOG_RANK=${LOG_RANK:-0}
export OMP_NUM_THREADS=${OMP_NUM_THREADS:-1}
export NCCL_P2P_DISABLE=${NCCL_P2P_DISABLE:-1}
export NCCL_SHM_LOCALITY=${NCCL_SHM_LOCALITY:-1}
export NCCL_BUFFSIZE=${NCCL_BUFFSIZE:-2097152}
MODULE=${MODULE:-"llama3"}
CONFIG=${CONFIG:-"llama3_debugmodel"}
COMM_MODE=${COMM_MODE:-""}
Expand Down
73 changes: 16 additions & 57 deletions torchtitan/components/metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,9 +41,7 @@ def __init__(self, device: str = f"{device_type}:0"):
self.device = torch.device(device) # device object
self.device_name = device_module.get_device_name(self.device)
self.device_index = device_module.current_device()
self.device_capacity = device_module.get_device_properties(
self.device
).total_memory
self.device_capacity = device_module.get_device_properties(self.device).total_memory
self.device_capacity_gib = self._to_gib(self.device_capacity)

device_module.reset_peak_memory_stats()
Expand Down Expand Up @@ -73,9 +71,7 @@ def get_peak_stats(self):
num_ooms = device_info.get("num_ooms", -1)

if num_retries > 0:
logger.warning(
f"{num_retries} {device_type.upper()} memory allocation retries."
)
logger.warning(f"{num_retries} {device_type.upper()} memory allocation retries.")
if num_ooms > 0:
logger.warning(f"{num_ooms} {device_type.upper()} OOM errors thrown.")

Expand Down Expand Up @@ -108,9 +104,7 @@ def log(self, metrics: dict[str, Any], step: int) -> None:
pass

def write_report(self, data: Any, step: int, name: str, output_type: str) -> None:
raise NotImplementedError(
f"{type(self).__name__} does not support write_report"
)
raise NotImplementedError(f"{type(self).__name__} does not support write_report")

def close(self) -> None:
pass
Expand Down Expand Up @@ -168,10 +162,7 @@ def __init__(
logger.info("WandB logging enabled")

def log(self, metrics: dict[str, Any], step: int) -> None:
wandb_metrics = {
(k if self.tag is None else f"{self.tag}/{k}"): v
for k, v in metrics.items()
}
wandb_metrics = {(k if self.tag is None else f"{self.tag}/{k}"): v for k, v in metrics.items()}
self.wandb.log(wandb_metrics, step=step)

def close(self) -> None:
Expand All @@ -193,9 +184,7 @@ def __init__(
raise ValueError("REPORTERV2_HOST must be set to use ReporterV2 logging")
training_id = os.getenv("REPORTERV2_TRAINING_ID")
if not training_id:
raise ValueError(
"REPORTERV2_TRAINING_ID must be set to use ReporterV2 logging"
)
raise ValueError("REPORTERV2_TRAINING_ID must be set to use ReporterV2 logging")
from reporterv2 import ReporterV2

reporter_config = dict(config_dict or {})
Expand All @@ -204,11 +193,7 @@ def __init__(
metrics_config = {}
self.save_freq = int(metrics_config.get("save_freq", 1))
model_spec = reporter_config.get("model_spec", {})
model_name = (
model_spec.get("name", "torchtitan")
if isinstance(model_spec, dict)
else "torchtitan"
)
model_name = model_spec.get("name", "torchtitan") if isinstance(model_spec, dict) else "torchtitan"
reporter_config["training_id"] = training_id
reporter_config["reporterv2_host"] = host
reporter_config.setdefault("trainer", model_name)
Expand All @@ -231,10 +216,6 @@ def log(self, metrics: dict[str, Any], step: int) -> None:
self.reporter.save_metrics()

def write_report(self, data: Any, step: int, name: str, output_type: str) -> None:
if output_type == "scalar":
self.reporter.buffer_metrics(step=step, epoch=step, metrics={name: data})
self.reporter.save_metrics()
return
self.reporter.write_report(data, step=step, name=name, output_type=output_type)

def close(self) -> None:
Expand Down Expand Up @@ -273,9 +254,7 @@ def close(self) -> None:
logger_instance.close()


def ensure_pp_loss_visible(
*, parallel_dims: ParallelDims, pp_schedule: str, color: Color | NoColor
) -> None:
def ensure_pp_loss_visible(*, parallel_dims: ParallelDims, pp_schedule: str, color: Color | NoColor) -> None:
"""
Ensures that the loss is visible on the console for pipeline-parallel training.

Expand Down Expand Up @@ -428,9 +407,7 @@ def __init__(
# used for colorful printing
self.color = utils.NoColor() if config.disable_color_printing else utils.Color()

self.gpu_peak_flops = utils.get_peak_flops(
self.device_memory_monitor.device_name
)
self.gpu_peak_flops = utils.get_peak_flops(self.device_memory_monitor.device_name)
self.ntokens_since_last_log = 0
self.data_loading_times = []
self.time_last_log = time.perf_counter()
Expand Down Expand Up @@ -469,23 +446,15 @@ def _build_metric_logger(
)

# Check if any logging backend is enabled
has_logging_enabled = (
config.enable_tensorboard
or config.enable_wandb
or config.enable_reporterv2
)
has_logging_enabled = config.enable_tensorboard or config.enable_wandb or config.enable_reporterv2

# Determine if this rank should log
should_log = has_logging_enabled
if (not config.save_for_all_ranks) and should_log:
metrics_rank = _get_metrics_rank(
parallel_dims=parallel_dims, pp_schedule=pp_schedule
)
metrics_rank = _get_metrics_rank(parallel_dims=parallel_dims, pp_schedule=pp_schedule)
should_log = torch.distributed.get_rank() == metrics_rank

logger.debug(
f"Logging decision: has_logging_enabled={has_logging_enabled}, should_log={should_log}"
)
logger.debug(f"Logging decision: has_logging_enabled={has_logging_enabled}, should_log={should_log}")

if not should_log:
logger.debug("Returning BaseLogger due to should_log=False")
Expand All @@ -505,9 +474,7 @@ def _build_metric_logger(
)

if config.save_for_all_ranks:
base_log_dir = os.path.join(
base_log_dir, f"rank_{torch.distributed.get_rank()}"
)
base_log_dir = os.path.join(base_log_dir, f"rank_{torch.distributed.get_rank()}")

# Create logger container
logger_container = LoggerContainer()
Expand All @@ -516,9 +483,7 @@ def _build_metric_logger(
if config.enable_wandb:
logger.debug("Attempting to create WandB logger")
try:
wandb_logger = WandBLogger(
base_log_dir, config_dict=config_dict, tag=tag
)
wandb_logger = WandBLogger(base_log_dir, config_dict=config_dict, tag=tag)
logger_container.add_logger(wandb_logger)
except Exception as e:
if "No module named 'wandb'" in str(e):
Expand Down Expand Up @@ -577,9 +542,7 @@ def log(
time_delta = time.perf_counter() - self.time_last_log

# tokens per second per device, abbreviated as tps
tps = self.ntokens_since_last_log / (
time_delta * self.parallel_dims.non_data_parallel_size
)
tps = self.ntokens_since_last_log / (time_delta * self.parallel_dims.non_data_parallel_size)
# model FLOPS utilization
# For its definition and calculation, please refer to the PaLM paper:
# https://arxiv.org/abs/2204.02311
Expand Down Expand Up @@ -639,17 +602,13 @@ def log(
self.time_last_log = time.perf_counter()
self.device_memory_monitor.reset_peak_stats()

def log_validation(
self, loss: float, step: int, extra_metrics: dict[str, Any] | None = None
):
def log_validation(self, loss: float, step: int, extra_metrics: dict[str, Any] | None = None):
time_delta = time.perf_counter() - self.time_last_log

device_mem_stats = self.device_memory_monitor.get_peak_stats()

# tokens per second per device, abbreviated as tps
tps = self.ntokens_since_last_log / (
time_delta * self.parallel_dims.non_data_parallel_size
)
tps = self.ntokens_since_last_log / (time_delta * self.parallel_dims.non_data_parallel_size)

metrics = {
"validation_metrics/loss": loss,
Expand Down
1 change: 1 addition & 0 deletions torchtitan/experiments/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
"autoparallel.llama3",
"autoparallel.local_map_deepseek_v3",
"path",
"rldriving",
"worldmodel",
"torchft.llama3",
"rl",
Expand Down
13 changes: 0 additions & 13 deletions torchtitan/experiments/path/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,16 +3,3 @@
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.

from .config_registry import model_registry
from .model import parallelize_path, PathModel
from .trainer import PathTrainer
from .validate import PathValidator

__all__ = [
"PathModel",
"PathTrainer",
"PathValidator",
"model_registry",
"parallelize_path",
]
Loading
Loading