Point-E 是 OpenAI 开源的文本生成三维点云模型,它能在几秒到几十秒内根据一句文字描述产出一个包含 1024 个点的三维点云。但在实际项目中,我们往往不是要一个点云,而是成千上万个,比如为下游模型准备训练集、做点云补全研究、构建三维检索数据库等。这时手动一个个输入提示词显然不现实,写一个可复用的批量生成脚本就成了刚需。本文将从环境搭建、脚本编写、性能加速到工程化细节,完整讲清楚如何用 Point-E 搭建一条自动化的点云数据生产线。

环境准备与模型加载
Point-E 依赖 PyTorch 运行,官方推荐使用 Python 3.8 以上版本。安装方式很简单,直接通过 pip 安装 point-e 包即可,同时建议安装 open3d 或 trimesh 用于点云的读取、可视化和格式转换。安装命令如下:
pip install point-e open3d trimesh
加载模型是批量脚本的第一步。Point-E 的推理分两个阶段:先用文本到图像的扩散模型生成一张合成图,再用图像到点云的模型把图转成点云。因此脚本中需要加载 base40M-text2pt 这类复合模型。为了在批量任务中避免反复加载模型浪费时间,务必把加载逻辑放在循环外部,只执行一次:
import torch
from point_e.util.pc_to_mesh import marching_cubes_mesh
from point_e.diffusion.configs import DIFFUSION_CONFIGS
from point_e.diffusion.sampler import PointCloudSampler
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 只加载一次模型,后续所有生成任务共用
sampler = PointCloudSampler(
device=device,
models=DIFFUSION_CONFIGS['base40M-text2pt']['models'],
auxiliaries=DIFFUSION_CONFIGS['base40M-text2pt']['auxiliaries'],
)
需要注意的是,模型加载到 GPU 大约占 2GB 显存,如果显存紧张可以退回 CPU 运行,但速度会慢一个数量级以上。另外首次运行会自动下载模型权重,建议提前手动下载并放到缓存目录,避免批量任务中途因网络问题中断。
编写批量生成核心逻辑
批量脚本的核心是一个提示词列表加一个生成循环。提示词可以从文本文件读取,也可以用模板组合生成。下面的例子演示了从 prompts.txt 逐行读取提示词,为每个提示词生成一个点云并保存为 PLY 文件:
import os
from point_e.util.plotting import plot_point_cloud
output_dir = './pointcloud_dataset'
os.makedirs(output_dir, exist_ok=True)
with open('prompts.txt', 'r', encoding='utf-8') as f:
prompts = [line.strip() for line in f if line.strip()]
for i, prompt in enumerate(prompts):
# 设置随机种子,保证结果可复现
torch.manual_seed(i)
samples = sampler.sample_batch(prompt=prompt, batch_size=1, guidance_scale=3.0)
pc = sampler.output_to_point_clouds(samples)[0]
# 保存为 PLY 格式
file_path = os.path.join(output_dir, f'pc_{i:05d}.ply')
pc.save(file_path)
print(f'[{i+1}/{len(prompts)}] 已保存: {file_path}')
这段代码有几个值得注意的细节。首先是 torch.manual_seed(i),为每个样本设定独立种子,这样批量结果既有多样性又完全可复现,做实验对比时非常重要。其次是文件命名用了五位数编号 pc_00001.ply,这样文件按名称排序就是生成顺序,方便后续索引。
其次,建议同时保存一份元数据文件,记录每个点云对应的提示词、随机种子和生成参数。这在数据集回溯和去重时非常有用,否则几千个文件生成完,很难分清哪个是哪个。可以用 JSON 或者 CSV 简单记录,例如把 {'id': i, 'prompt': prompt, 'seed': i, 'guidance_scale': 3.0} 逐行写入一个 metadata.jsonl 文件。
性能优化:批处理与多进程加速
逐条生成是最直观但效率最低的方式。Point-E 的采样器原生支持 batch,一次前向可以同时生成多个点云,GPU 利用率会显著提升。如果显存允许(8GB 以上一般可以跑 batch_size 为 4 到 8),推荐按批次组织提示词:
batch_size = 4
for start in range(0, len(prompts), batch_size):
batch_prompts = prompts[start:start + batch_size]
samples = sampler.sample_batch(
prompt=batch_prompts,
batch_size=len(batch_prompts),
guidance_scale=3.0,
)
clouds = sampler.output_to_point_clouds(samples)
for j, pc in enumerate(clouds):
idx = start + j
pc.save(os.path.join(output_dir, f'pc_{idx:05d}.ply'))
除了批处理,还可以在采样步数上做权衡。默认配置走 64 步扩散,如果对质量要求不高,把 diffusion_steps 降到 32 甚至 16,速度几乎线性提升,而生成质量只略有下降,特别适合快速预览或生成草稿级数据。另外 guidance_scale 控制文本贴合程度,调高会更忠实于描述但多样性下降,批量生成数据集时通常设在 3.0 左右比较均衡。
如果要进一步榨干硬件,可以考虑多进程方案:每个进程独立加载一份模型,分别处理提示词列表的不同分片。但这种方式要求显存够大或者搭配多张卡,否则多个进程会争抢显存导致 OOM。更经济的做法是单进程模型常驻,用生产者消费者队列把生成和保存解耦——生成线程专注 GPU 计算,保存操作交给另一个线程做 IO,两者重叠执行,整体吞吐能提升两到三成。
工程化细节:断点续跑与失败重试
大规模生成任务跑几个小时是常态,中途断电、显存溢出、个别提示词触发异常都在所难免,所以脚本必须具备断点续跑能力。实现思路很简单:启动时扫描输出目录,把已经存在的文件对应的任务跳过:
def get_pending_tasks(prompts, output_dir):
tasks = []
for i, prompt in enumerate(prompts):
if not os.path.exists(os.path.join(output_dir, f'pc_{i:05d}.ply')):
tasks.append((i, prompt))
return tasks
for i, prompt in get_pending_tasks(prompts, output_dir):
try:
torch.manual_seed(i)
samples = sampler.sample_batch(prompt=prompt, batch_size=1)
pc = sampler.output_to_point_clouds(samples)[0]
pc.save(os.path.join(output_dir, f'pc_{i:05d}.ply'))
except Exception as e:
print(f'任务 {i} 失败: {e},已记录,继续下一个')
with open('failed_tasks.log', 'a') as log:
log.write(f'{i}\t{prompt}\t{e}\n')
失败重试的粒度也要控制好。单个任务失败不应该终止整个流程,用 try except 捕获后记录到日志文件,跑完一轮后可以针对失败列表再跑一遍。此外,长时间运行还要注意 GPU 显存碎片问题,PyTorch 可以定期调用 torch.cuda.empty_cache() 释放缓存,或者每隔几百个任务重启进程配合外部调度脚本(比如用 shell 脚本循环拉起 Python 脚本,靠断点续跑机制衔接进度)。
最后在数据格式上,PLY 是最通用的选择,Open3D、MeshLab、CloudCompare 都能直接打开。如果后续要喂给深度学习模型,也可以额外导出 NumPy 的 .npy 文件,直接用 pc.coords 拿到形状为 (1024, 3) 的坐标数组,配合颜色通道 pc.channels 使用,加载速度比解析 PLY 快得多。
总的来说,用 Point-E 批量生成点云数据集并不复杂,关键在于把模型加载、批处理采样、断点续跑、元数据记录这几个工程环节做扎实。只要脚本写得稳,一张消费级显卡一晚上就能产出上万个点云样本,足以支撑大多数研究和原型验证场景的数据需求。