目录
mmyy5213个月前3次提交

Jittor PCT for ModelNet40 Classification

本项目使用 Jittor 深度学习框架实现 PCT(Point Cloud Transformer)点云分类模型,用于 ModelNet40 三维形状分类任务。模型输入为每个样本 2048 个三维点,输出测试集中每个样本的 40 类分类结果,并生成比赛要求的 result.json

任务说明

  • 框架:Jittor
  • 模型:PCT baseline
  • 数据集:ModelNet40 预处理点云数据
  • 训练集:9843 个样本
  • 测试集:2468 个样本
  • 点数:每个样本 2048 个点,训练时可通过 --n_points 随机采样
  • 类别数:40
  • 评测指标:Accuracy
  • 通过线:测试集 Accuracy >= 0.80

提交文件格式为:

result.zip
└── result.json

result.json 为字典格式,key 是测试样本编号字符串,value 是预测类别编号:

{
  "0": 4,
  "1": 35,
  "2": 10
}

环境安装

推荐使用 Python 3.9 和已有的 jittor conda 环境:

conda activate jittor

或直接使用环境中的 Python:

/root/anaconda3/envs/jittor/bin/python -m jittor.test.test_example

如果需要从零安装,可参考 Jittor 官方安装方式:

pip install jittor
python -m jittor.test.test_example

GPU 训练需要正确安装 CUDA 与 cudnn-dev。若出现 libcudnn_ops_infer.so not found,需要确认 CUDNN 库位于 Jittor 根据 nvcc_path 找到的 CUDA 目录下,例如:

ls /usr/local/cuda/lib64/libcudnn_ops_infer.so
ls /usr/local/cuda/include/cudnn.h

数据准备

将比赛提供的 ModelNet40 预处理数据解压到项目根目录的 data/ 下,目录结构如下:

jittor-pct/
├── pct.py
└── data/
    ├── train_points.npy
    ├── train_labels.npy
    ├── test_points.npy
    └── categories.txt

数据文件说明:

  • data/train_points.npy:训练点云,shape 为 (9843, 2048, 3)
  • data/train_labels.npy:训练标签,shape 为 (9843,)
  • data/test_points.npy:测试点云,shape 为 (2468, 2048, 3)
  • data/categories.txt:40 个类别名称

开源仓库中不建议提交 .npy 数据文件和 data.zip,只保留数据准备说明。

训练

本项目使用的训练命令为:

python pct.py --epochs 200 --batch_size 32 --n_points 1024 --lr 0.01

如果使用指定 conda 环境中的 Python,可运行:

/root/anaconda3/envs/jittor/bin/python pct.py --epochs 200 --batch_size 32 --n_points 1024 --lr 0.01

主要参数说明:

  • --data_dir:数据目录,默认 ./data
  • --n_points:每个样本训练时采样的点数,当前使用 1024
  • --batch_size:batch size,当前使用 32
  • --epochs:训练轮数,当前使用 200
  • --lr:初始学习率,当前使用 0.01
  • --seed:随机种子,默认 42

训练完成后会生成:

pct_model.pkl
result.json

推理与提交

pct.py 在训练结束后会自动使用训练好的模型对测试集进行预测,并保存 result.json

生成提交压缩包:

zip result.zip result.json

提交前建议检查预测数量:

python - <<'CHECK_RESULT'
import json

with open("result.json", "r") as f:
    result = json.load(f)

assert len(result) == 2468, len(result)
assert all(isinstance(k, str) for k in result.keys())
assert all(isinstance(v, int) and 0 <= v < 40 for v in result.values())
print("result.json format ok")
CHECK_RESULT

结果说明

线上评测程序会读取 result.json,与隐藏测试标签比较并计算 Accuracy:

Accuracy = 正确预测样本数 / 测试样本总数

PCT baseline 的测试准确率通常在 80% 通过线附近波动。由于测试集标签不公开,本仓库只提供训练与提交文件生成流程,最终成绩以线上评测结果为准。

复现说明

建议固定随机种子并记录运行命令:

python pct.py --epochs 200 --batch_size 32 --n_points 1024 --lr 0.01 --seed 42

推荐保存训练日志:

python pct.py --epochs 200 --batch_size 32 --n_points 1024 --lr 0.01 --seed 42 2>&1 | tee train.log
关于
35.0 KB
邀请码
    Gitlink(确实开源)
  • 加入我们
  • 官网邮箱:gitlink@ccf.org.cn
  • QQ群
  • QQ群
  • 公众号
  • 公众号

版权所有:中国计算机学会技术支持:开源发展技术委员会
京ICP备13000930号-9 京公网安备 11010802047560号