import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
import os

# 设置中文字体（如果有的话）
plt.rcParams['font.sans-serif'] = ['DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False

# 模拟一个3层网络，中间层权重更新前后的输入分布变化
np.random.seed(42)

# 原始数据：一批样本（比如batch=5）
# 假设是某个中间层的输入
batch_data = np.array([1.0, 2.0, 3.0, 4.0, 5.0])

# 权重更新前：第2层输出（经过线性变换后）
# 比如 weight=1, bias=0 → 输出=[1,2,3,4,5]
pre_update = batch_data * 1.0 + 0.0

# 权重更新后：weight变成了3，bias变成了2 → 输出=[5,8,11,14,17]
post_update = batch_data * 3.0 + 2.0

fig, axes = plt.subplots(2, 3, figsize=(15, 9))

# ============ 第1行：没有 BN ============
ax = axes[0, 0]
ax.bar(range(5), pre_update, color='steelblue', edgecolor='navy')
ax.set_title('Layer 2 Output (Before Weight Update)', fontsize=12, fontweight='bold')
ax.set_ylabel('Value')
ax.set_xticks(range(5))
ax.set_xticklabels([f'Sample {i+1}' for i in range(5)])
ax.set_ylim(0, 20)
for i, v in enumerate(pre_update):
    ax.text(i, v+0.3, f'{v:.1f}', ha='center', fontsize=10)
mean, std = pre_update.mean(), pre_update.std()
ax.text(0.5, 18, f'Mean={mean:.1f}, Std={std:.2f}', fontsize=10, 
        bbox=dict(boxstyle='round', facecolor='yellow', alpha=0.7))

ax = axes[0, 1]
ax.bar(range(5), post_update, color='coral', edgecolor='darkred')
ax.set_title('Layer 2 Output (After Weight Update)', fontsize=12, fontweight='bold')
ax.set_ylabel('Value')
ax.set_xticks(range(5))
ax.set_xticklabels([f'Sample {i+1}' for i in range(5)])
ax.set_ylim(0, 20)
for i, v in enumerate(post_update):
    ax.text(i, v+0.3, f'{v:.1f}', ha='center', fontsize=10)
mean, std = post_update.mean(), post_update.std()
ax.text(0.5, 18, f'Mean={mean:.1f}, Std={std:.2f}', fontsize=10,
        bbox=dict(boxstyle='round', facecolor='yellow', alpha=0.7))

ax = axes[0, 2]
# 画分布变化：两个分布并排
data = [pre_update, post_update]
bp = ax.boxplot(data, labels=['Before Update', 'After Update'], patch_artist=True)
bp['boxes'][0].set_facecolor('steelblue')
bp['boxes'][1].set_facecolor('coral')
ax.set_title('Layer 3 Sees: Distribution Changed!', fontsize=12, fontweight='bold')
ax.set_ylabel('Value')
ax.set_ylim(0, 20)
ax.axhline(y=5, color='gray', linestyle='--', alpha=0.5, label='Layer 3 expected range')
ax.text(1.5, 17, 'Layer 3 must re-learn!', fontsize=11, color='red', fontweight='bold',
        ha='center', bbox=dict(boxstyle='round', facecolor='lightyellow', alpha=0.8))

# ============ 第2行：有 BN ============
# BN 标准化：减去均值，除以标准差
pre_bn = (pre_update - pre_update.mean()) / (pre_update.std() + 1e-5)
post_bn = (post_update - post_update.mean()) / (post_update.std() + 1e-5)

ax = axes[1, 0]
ax.bar(range(5), pre_bn, color='steelblue', edgecolor='navy')
ax.set_title('After BN (Before Update)', fontsize=12, fontweight='bold')
ax.set_ylabel('Normalized Value')
ax.set_xticks(range(5))
ax.set_xticklabels([f'Sample {i+1}' for i in range(5)])
ax.set_ylim(-2, 2)
for i, v in enumerate(pre_bn):
    ax.text(i, v+0.05, f'{v:.2f}', ha='center', fontsize=10)
mean, std = pre_bn.mean(), pre_bn.std()
ax.text(0.5, 1.7, f'Mean={mean:.2f}, Std={std:.2f}', fontsize=10,
        bbox=dict(boxstyle='round', facecolor='lightgreen', alpha=0.7))

ax = axes[1, 1]
ax.bar(range(5), post_bn, color='coral', edgecolor='darkred')
ax.set_title('After BN (After Update)', fontsize=12, fontweight='bold')
ax.set_ylabel('Normalized Value')
ax.set_xticks(range(5))
ax.set_xticklabels([f'Sample {i+1}' for i in range(5)])
ax.set_ylim(-2, 2)
for i, v in enumerate(post_bn):
    ax.text(i, v+0.05, f'{v:.2f}', ha='center', fontsize=10)
mean, std = post_bn.mean(), post_bn.std()
ax.text(0.5, 1.7, f'Mean={mean:.2f}, Std={std:.2f}', fontsize=10,
        bbox=dict(boxstyle='round', facecolor='lightgreen', alpha=0.7))

ax = axes[1, 2]
data_bn = [pre_bn, post_bn]
bp = ax.boxplot(data_bn, labels=['Before Update', 'After Update'], patch_artist=True)
bp['boxes'][0].set_facecolor('steelblue')
bp['boxes'][1].set_facecolor('coral')
ax.set_title('Layer 3 Sees: Same Distribution!', fontsize=12, fontweight='bold')
ax.set_ylabel('Normalized Value')
ax.set_ylim(-2, 2)
ax.axhline(y=0, color='gray', linestyle='--', alpha=0.5)
ax.text(1.5, 1.5, 'Layer 3 is happy!', fontsize=11, color='green', fontweight='bold',
        ha='center', bbox=dict(boxstyle='round', facecolor='lightgreen', alpha=0.8))

# 总标题
fig.suptitle('Why BatchNorm Works: A Concrete Example', fontsize=16, fontweight='bold', y=1.02)

plt.tight_layout()

# 保存到 code 目录
out_path = os.path.join(os.path.dirname(__file__), 'batchnorm_example.png')
plt.savefig(out_path, dpi=150, bbox_inches='tight')
print(f"Saved: {out_path}")

# 同时保存到 copyparty
copyparty_path = '/root/copyparty-files/deeplearning/code/batchnorm_example.png'
plt.savefig(copyparty_path, dpi=150, bbox_inches='tight')
print(f"Synced to copyparty: {copyparty_path}")

plt.close()
