You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何制作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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 07:02:02