# Visualize the bit layout of each format
fig, axes = plt.subplots(3, 1, figsize=(14, 5))
formats = [
('FP32 (32 bits)', 1, 8, 23, '#94A3B8'),
('FP16 (16 bits)', 1, 5, 10, '#FF6B6B'),
('BF16 (16 bits)', 1, 8, 7, '#4ECDC4'),
]
for ax, (name, sign_bits, exp_bits, mant_bits, color) in zip(axes, formats):
total = sign_bits + exp_bits + mant_bits
x = 0
# Sign bit
rect = mpatches.FancyBboxPatch((x, 0), sign_bits, 1, boxstyle='round,pad=0.02',
facecolor='#C4B5FD', edgecolor='#333')
ax.add_patch(rect)
ax.text(x + sign_bits/2, 0.5, 'S', ha='center', va='center', fontweight='bold', fontsize=10)
x += sign_bits
# Exponent
rect = mpatches.FancyBboxPatch((x, 0), exp_bits, 1, boxstyle='round,pad=0.02',
facecolor='#FFE66D', edgecolor='#333')
ax.add_patch(rect)
ax.text(x + exp_bits/2, 0.5, f'Exponent ({exp_bits} bits)', ha='center', va='center',
fontweight='bold', fontsize=10)
x += exp_bits
# Mantissa
rect = mpatches.FancyBboxPatch((x, 0), mant_bits, 1, boxstyle='round,pad=0.02',
facecolor=color, edgecolor='#333')
ax.add_patch(rect)
ax.text(x + mant_bits/2, 0.5, f'Mantissa ({mant_bits} bits)', ha='center', va='center',
fontweight='bold', fontsize=10)
ax.set_xlim(-0.5, 32.5)
ax.set_ylim(-0.2, 1.5)
ax.set_ylabel(name, fontsize=11, fontweight='bold')
ax.set_xticks([])
ax.set_yticks([])
ax.spines['top'].set_visible(False)
ax.spines['right'].set_visible(False)
ax.spines['bottom'].set_visible(False)
ax.spines['left'].set_visible(False)
axes[0].set_title('Floating Point Bit Layouts', fontsize=14, fontweight='bold', pad=10)
plt.tight_layout()
plt.show()
print('Key: more exponent bits = larger range, more mantissa bits = more precision')
print('BF16 has FP32\'s range (8-bit exponent) but coarser precision (7-bit mantissa)')