Matplotlib添加大量图像缩略图内存不足,求栅格化解决办法
解决Matplotlib批量添加图像时的内存不足问题
我尝试制作类似UMAP Explorer的大型2D嵌入可视化图像,其中包含大量嵌入内容:
使用的原始代码如下:
def getImage(path, width=100, zoom=1.0): image = plt.imread(path) image = imutils.resize(image, width=width) return OffsetImage(image, zoom=zoom) fig, ax = plt.subplots(dpi=100) ax.scatter( embeddings[:, 0], embeddings[:, 1]) for x, y, path in zip(embeddings[:, 0], embeddings[:, 1], image_paths): ab = AnnotationBbox(getImage(path, width=1024, zoom=0.005), (x, y), frameon=False) ax.add_artist(ab)
当图像数量超过可用内存时,代码因内存耗尽无法运行,推测是Matplotlib未对图像进行栅格化,而是持续保留所有图像的原始数据引用导致内存占用过高。以下是几种可行的解决方法:
方法1:提前栅格化并缓存小尺寸图像
不要每次循环都加载原始大尺寸图像,先统一将所有图像预处理为小尺寸位图并缓存,直接复用缓存后的图像数据,大幅降低内存占用:
import numpy as np import matplotlib.pyplot as plt from matplotlib.offsetbox import OffsetImage, AnnotationBbox import imutils # 预处理所有图像为统一小尺寸,缓存为uint8格式数组 def preprocess_images(image_paths, target_width=32): cached_images = [] for path in image_paths: img = plt.imread(path) img = imutils.resize(img, width=target_width) # 转换为uint8格式减少内存占用(浮点数格式内存是其4倍) cached_images.append(img.astype(np.uint8)) return cached_images def get_cached_image(img_array, zoom=1.0): return OffsetImage(img_array, zoom=zoom) # 主流程 fig, ax = plt.subplots(dpi=100) ax.scatter(embeddings[:, 0], embeddings[:, 1]) # 预处理所有图像 cached_imgs = preprocess_images(image_paths, target_width=32) for x, y, img in zip(embeddings[:, 0], embeddings[:, 1], cached_imgs): ab = AnnotationBbox(get_cached_image(img, zoom=1.0), (x, y), frameon=False) ax.add_artist(ab) plt.savefig("umap_visualization.png", dpi=150)
方法2:给AnnotationBbox启用栅格化
直接给AnnotationBbox添加rasterized=True参数,强制Matplotlib在渲染时栅格化图像元素,避免保留矢量数据:
# 修改循环内的AnnotationBbox创建代码 ab = AnnotationBbox(getImage(path, width=1024, zoom=0.005), (x, y), frameon=False, rasterized=True)
注意:该方法对PNG等位图格式保存效果更明显,若保存为PDF等矢量格式,需确保栅格化参数生效。
方法3:分批次渲染后合成图像
若图像数量极大(数万级),超过单进程内存上限,可分区域渲染图像图层,再合成最终大图:
import numpy as np import matplotlib.pyplot as plt from matplotlib.offsetbox import OffsetImage, AnnotationBbox import imutils from PIL import Image # 先保存基础散点图作为底图 fig, ax = plt.subplots(dpi=100) ax.scatter(embeddings[:, 0], embeddings[:, 1]) plt.savefig("base_map.png", dpi=100, bbox_inches='tight') plt.close() # 按x轴划分区域,分批次渲染 x_min, x_max = embeddings[:,0].min(), embeddings[:,0].max() x_bins = np.linspace(x_min, x_max, 4) # 分为3个区域 base_img = Image.open("base_map.png") for i in range(3): # 筛选当前区域内的嵌入点和图像路径 mask = (embeddings[:,0] >= x_bins[i]) & (embeddings[:,0] < x_bins[i+1]) batch_x = embeddings[mask, 0] batch_y = embeddings[mask, 1] batch_paths = [path for idx, path in enumerate(image_paths) if mask[idx]] # 创建临时画布渲染当前批次图像(隐藏坐标轴) fig, ax = plt.subplots(dpi=100) ax.set_xlim(x_min, x_max) ax.set_ylim(embeddings[:,1].min(), embeddings[:,1].max()) ax.axis('off') for x, y, path in zip(batch_x, batch_y, batch_paths): img = plt.imread(path) img = imutils.resize(img, width=32) oi = OffsetImage(img, zoom=1.0) ab = AnnotationBbox(oi, (x, y), frameon=False) ax.add_artist(ab) # 保存透明背景的图层 plt.savefig(f"layer_{i}.png", dpi=100, bbox_inches='tight', transparent=True) plt.close() # 将图层合成到底图 layer_img = Image.open(f"layer_{i}.png") base_img.paste(layer_img, (0,0), layer_img) # 保存最终合成图 base_img.save("final_umap.png")
额外优化建议
- 调整
target_width和zoom参数,平衡图像清晰度与内存占用,尺寸越小内存消耗越低 - 若存在重复的图像路径,提前做去重缓存,避免重复加载同一图像
- 尽量使用
uint8格式存储图像数组,相比浮点数格式可节省75%内存
内容的提问来源于stack exchange,提问作者Cypher
相关产品推荐
相关产品推荐

