如何制作MNIST数据集NMF组件随总组件数变化的子图动画?
实现NMF组件数递减的动画展示
没问题,要实现这个组件变化的动画其实很简单,咱们用matplotlib的FuncAnimation就能搞定,我给你整理了完整的实现步骤和代码:
1. 加载并预处理MNIST数据(和你之前的代码一致)
先把数据加载成NumPy数组,确保后续处理没问题:
import numpy as np from mnist import MNIST from sklearn.decomposition import NMF import matplotlib.pyplot as plt from matplotlib.animation import FuncAnimation # 加载MNIST数据 mndata = MNIST('./data') images_train, labels_train = mndata.load_training() X_train = np.array(images_train).astype('float64')
2. 定义生成NMF组件的复用函数
写一个通用函数,输入组件数量,返回训练好的NMF模型的前15个组件:
def get_nmf_components(n_components): # 初始化并训练NMF模型,增加迭代次数确保收敛 nmf = NMF(n_components=n_components, random_state=0, max_iter=500) nmf.fit(X_train) # 返回前15个组件(咱们的组件数范围是15到100,不用担心数量不足) return nmf.components_[:15]
3. 设置动画的核心参数
确定组件数的变化序列:从100递减到15,步长为5:
# 组件数序列:100, 95, ..., 15 n_components_list = list(range(100, 14, -5)) frame_count = len(n_components_list)
4. 初始化绘图布局
创建和你之前一致的3行5列子图布局:
fig, axes = plt.subplots(3, 5, figsize=(12, 7), subplot_kw={'xticks': (), 'yticks': ()}) axes = axes.ravel() # 把二维子图数组展平成一维,方便遍历 # 初始化第一帧的图像(用100组件的情况) initial_components = get_nmf_components(100) imgs = [] for i, ax in enumerate(axes): img = ax.imshow(initial_components[i].reshape(28,28), cmap=plt.cm.binary) imgs.append(img) ax.set_title(f"{i+1}. component") # 添加总标题,显示当前的总组件数 title = fig.suptitle(f"NMF Components (Total: {n_components_list[0]})", fontsize=14)
5. 定义动画更新函数
这个函数会在每一帧被调用,负责更新所有子图的图像和总标题:
def update(frame): current_n = n_components_list[frame] # 获取当前组件数对应的前15个组件 components = get_nmf_components(current_n) # 更新每个子图的图像 for i, (img, component) in enumerate(zip(imgs, components)): img.set_data(component.reshape(28,28)) # 更新总标题 title.set_text(f"NMF Components (Total: {current_n})") # 返回所有需要更新的元素(图像和标题) return imgs + [title]
6. 创建并展示/保存动画
用FuncAnimation生成动画,你可以选择直接展示,或者保存成视频文件:
# 创建动画,interval是每帧之间的间隔(毫秒),这里设为1000即1秒切换一次 ani = FuncAnimation(fig, update, frames=frame_count, interval=1000, blit=True) # 如果你想保存动画为MP4,需要先安装ffmpeg,然后取消下面的注释 # ani.save('nmf_components_animation.mp4', writer='ffmpeg', dpi=100) # 展示动画 plt.show()
优化小技巧
因为每次更新都要重新训练NMF,这个过程可能会有点慢。你可以提前预计算所有组件数对应的结果,然后在动画里直接读取,这样播放会流畅很多:
# 预计算所有组件数对应的前15个组件 precomputed_components = [] for n in n_components_list: print(f"Training NMF with {n} components...") precomputed_components.append(get_nmf_components(n)) # 修改update函数,直接读取预计算好的结果 def update(frame): current_n = n_components_list[frame] components = precomputed_components[frame] for i, (img, component) in enumerate(zip(imgs, components)): img.set_data(component.reshape(28,28)) title.set_text(f"NMF Components (Total: {current_n})") return imgs + [title]
内容的提问来源于stack exchange,提问作者Fallen Apart
相关产品推荐
相关产品推荐

