-
Notifications
You must be signed in to change notification settings - Fork 32
Expand file tree
/
Copy pathvisualize_logs.py
More file actions
133 lines (98 loc) · 4.27 KB
/
Copy pathvisualize_logs.py
File metadata and controls
133 lines (98 loc) · 4.27 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
import argparse
import json
from pathlib import Path
import pandas as pd
import matplotlib.pyplot as plt
# Load a user log file and convert it to pandas DataFrame
def load_user_log(path: Path) -> pd.DataFrame:
with path.open("r", encoding="utf-8") as f:
data = json.load(f)
logs = data.get("logs", [])
if not logs:
return pd.DataFrame()
df = pd.DataFrame(logs)
user_info = data.get("user", {})
run_info = data.get("run", {})
task_info = data.get("task", {})
df["user_index"] = user_info.get("index")
df["client_id"] = user_info.get("clientId")
df["test_id"] = run_info.get("testID")
df["task_id"] = run_info.get("taskID", task_info.get("id"))
df["step"] = range(len(df))
return df
# Merge logs from all clients into a single pd.DataFrame
def load_all_user_logs(log_dir: Path) -> pd.DataFrame:
files = sorted(log_dir.glob("*_local_log.json"))
dfs = [load_user_log(f) for f in files]
dfs = [df for df in dfs if not df.empty]
return pd.concat(dfs, ignore_index=True)
# Ensure that the specified columns are converted into numeric values
def ensure_numeric(df: pd.DataFrame, columns):
for col in columns:
if col in df.columns:
df[col] = pd.to_numeric(df[col], errors="coerce")
return df
# plotting functions
# Draw a line plot for one metric across all users
def plot_metric_per_user(df, metric, output_path, title, ylabel):
plt.figure(figsize=(10, 6))
for user, g in df.groupby("user_index"):
plt.plot(g["step"], g[metric], label=f"user {user}")
plt.xlabel("Step")
plt.ylabel(ylabel)
plt.title(title)
plt.legend()
plt.tight_layout()
plt.savefig(output_path)
plt.close()
# Plot mean value of a metric with std
def plot_mean_std(df, metric, output_path, title, ylabel):
summary = df.groupby("step")[metric].agg(["mean", "std"]).reset_index()
# Handle edge cases where only one client's value is available for a step (std becomes NaN)
summary["std"] = summary["std"].fillna(0)
plt.figure(figsize=(10, 6))
plt.plot(summary["step"], summary["mean"])
plt.fill_between(summary["step"], summary["mean"]-summary["std"], summary["mean"]+summary["std"], alpha=0.2)
plt.xlabel("Step")
plt.ylabel(ylabel)
plt.title(title)
plt.tight_layout()
plt.savefig(output_path)
plt.close()
print(f"Plot saved: {output_path}")
# Display multiple plots in a single figure
def plot_dashboard(df, output_path):
metrics = [
("trainingLoss", "Training Loss"), # metric and title for each subplot
("validationLoss", "Validation Loss"),
("trainingAccuracy", "Training Accuracy"),
("validationAccuracy", "Validation Accuracy")
]
fig, axes = plt.subplots(2, 2, figsize=(12, 8))
for ax, (metric, title) in zip(axes.flatten(), metrics):
for user, g in df.groupby("user_index"):
ax.plot(g["step"], g[metric])
ax.set_title(title)
ax.set_xlabel("Step")
plt.tight_layout()
plt.savefig(output_path)
plt.close()
def main():
parser = argparse.ArgumentParser()
parser.add_argument("log_dir", type=str, help="Path to log directory for visualization")
args = parser.parse_args()
log_dir = Path(args.log_dir)
df = load_all_user_logs(log_dir)
# Ensure numeric values for columns used in visualization
df = ensure_numeric(df, ["trainingLoss", "validationLoss", "trainingAccuracy", "validationAccuracy", "epochTime", "peakMemory"])
# Per-client plots
plot_metric_per_user(df, "trainingLoss", log_dir / "training_loss.png", "Training Loss", "Loss")
plot_metric_per_user(df, "trainingAccuracy", log_dir / "training_acc.png", "Training Accuracy", "Accuracy")
plot_metric_per_user(df, "validationLoss", log_dir / "validation_loss.png", "Validation Loss", "Loss")
plot_metric_per_user(df, "validationAccuracy", log_dir / "validation_acc.png", "Validation Accuracy", "Accuracy")
# Mean loss and accuracy plots
plot_mean_std(df, "validationLoss", log_dir / "mean_validation_loss.png", "Mean Validation Loss", "Loss")
plot_mean_std(df, "validationAccuracy", log_dir / "mean_validation_acc.png", "Mean Validation Accuracy", "Accuracy")
plot_dashboard(df, log_dir / "dashboard.png")
if __name__ == "__main__":
main()