Skip to content
Open
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
60 changes: 35 additions & 25 deletions tests/dlio_ai_logging_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
import pytest
import os
import glob
import re
from datetime import datetime
from collections import Counter

Expand Down Expand Up @@ -97,6 +98,14 @@ def check_ai_events(path):
counter["epoch"] += 1
return counter

def get_trace_files(storage_root):
paths = sorted(
path for path in glob.glob(os.path.join(storage_root, "*.pfw"))
if os.path.getsize(path) > 0
)
assert paths, f"No nonempty pfw files found in {storage_root}"
return paths

def get_rank_trace_files(all_paths, num_procs):
"""
Find main trace files for each MPI rank.
Expand All @@ -108,17 +117,17 @@ def get_rank_trace_files(all_paths, num_procs):
Returns:
Dictionary mapping rank number to trace file path
"""
# Filter to main trace files only (exclude worker traces like trace-{hash}-app.pfw)
main_traces = [p for p in all_paths if "-of-" in p and "-app.pfw" not in p]

rank_traces = {}
for rank in range(num_procs):
# Match pattern: trace-{rank}-of-{num_procs}.pfw
matching = [p for p in main_traces if f"trace-{rank}-of-{num_procs}.pfw" in p]
if matching:
rank_traces[rank] = matching[0]
else:
print(f"WARNING: No main trace file found for rank {rank}")
# Newer DFTracer appends a hash and -app to the main trace name.
# Worker traces include an additional process ID after the rank.
pattern = re.compile(rf"trace-{rank}-of-{num_procs}(?:\.pfw|-[0-9a-f]+-app\.pfw)")
matching = [path for path in all_paths if pattern.fullmatch(os.path.basename(path))]
assert len(matching) == 1, (
f"Expected one main trace for rank {rank}, found {matching}; "
f"available traces: {all_paths}"
)
rank_traces[rank] = matching[0]

return rank_traces

Expand Down Expand Up @@ -154,9 +163,7 @@ def test_ai_logging_train(setup_test_env, framework, num_data, batch_size):
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))

assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Aggregate item and preprocess counts globally
global_item_count = 0
Expand Down Expand Up @@ -228,8 +235,7 @@ def test_ai_logging_train_with_step(setup_test_env, framework, step, read_thread
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))
assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Aggregate item and preprocess counts globally
global_item_count = 0
Expand All @@ -252,7 +258,9 @@ def test_ai_logging_train_with_step(setup_test_env, framework, step, read_thread
assert count["epoch"] == num_epochs
assert count["train"] == num_epochs
assert count["eval"] == 0
assert count["fetch_iter"] == num_epochs * step
# A step limit can leave the iterator open. Recent DFTracer versions
# also record the final probe that closes it at each epoch boundary.
assert num_epochs * step <= count["fetch_iter"] <= num_epochs * (step + 1)
assert count["compute"] == num_epochs * step

assert count["ckpt_capture"] == 0
Expand Down Expand Up @@ -292,8 +300,7 @@ def test_ai_logging_with_eval(setup_test_env, framework):
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))
assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Aggregate item and preprocess counts globally
global_item_count = 0
Expand Down Expand Up @@ -361,8 +368,7 @@ def test_ai_logging_with_reader(setup_test_env, framework, fmt):
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))
assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Aggregate item and preprocess counts globally
global_item_count = 0
Expand All @@ -385,8 +391,14 @@ def test_ai_logging_with_reader(setup_test_env, framework, fmt):
assert count["epoch"] == num_epochs
assert count["train"] == num_epochs
assert count["eval"] == num_epochs
assert count["fetch_iter"] == 2 * num_epochs * (num_data_pp // batch_size)
assert count["compute"] == 2 * num_epochs * (num_data_pp // batch_size)
expected_iters = 2 * num_epochs * (num_data_pp // batch_size)
if fmt == "tfrecord":
# DFTracer records the end-of-sequence fetch for each train/eval
# iterator in every epoch.
assert count["fetch_iter"] == expected_iters + 2 * num_epochs
else:
assert count["fetch_iter"] == expected_iters
assert count["compute"] == expected_iters

assert count["ckpt_capture"] == 0
assert count["ckpt_restart"] == 0
Expand Down Expand Up @@ -449,8 +461,7 @@ def test_ai_logging_train_with_checkpoint(setup_test_env, framework, epoch_per_c
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))
assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Aggregate item and preprocess counts globally
global_item_count = 0
Expand Down Expand Up @@ -529,8 +540,7 @@ def test_ai_logging_checkpoint_only(setup_test_env, framework, num_checkpoint_wr
# Run benchmark in MPI subprocess
run_mpi_benchmark(overrides, num_procs=NUM_PROCS)

paths = glob.glob(os.path.join(storage_root, "*.pfw"))
assert len(paths) > 0, "No pfw files found"
paths = get_trace_files(storage_root)

# Get main trace files for each rank
rank_traces = get_rank_trace_files(paths, NUM_PROCS)
Expand Down
Loading