text2CAD模型复刻实践记录
0x00 论文概述
0x00.1 生成流程
由冷启动初步训练SFT模型,随后从几何体角度使用CoT强化学习训练模型。强化学习涉及奖励函数(相对优势\(\hat{A}\))与数据集构建部分,奖励函数由倒角距离CD,第二个损失函数有关。
0x00.2 主要专业术语
- GRPO:群体奖励策略优化(一个模型得到多个答案选最优)
- 几何奖励\(R_i^{geo}\):模型生成的几何体与答案几何体是否一致(由倒角距离CD衡量)
- 格式奖励\(R_i^{fmt}\):生成代码语法是否正确
- CoT:思维链(从策略规划到代码生成)
- SFT:监督微调(照着答案改)
- 命令序列S:模型模仿人类构建模型的思考过程
- L:自然语言提示词
- \(C_{gt}\) :正确Cadquery代码(一般以token方式处理)
- 损失函数:由于模型生成产生的与答案几何体的偏差。一般损失越小,答案越优。
- \(E_{(L,C_{gt})\sim D}\):期望
- 冷启动:第一轮学习前对模型进行优质数据的监督微调
- 倒角距离CD:双向找最近点。对于预测点阵与答案点阵中的每一个点,都找到另一个点阵相同坐标附近的最近点,把所有距离求和取平均(或者别的方法)
- \(||x-y||^2_2\):点x与点y之间的欧几里得距离
- \(P\):预测点阵
- \(Q\):答案点阵
- \(L_{cot}\):规划流程
- \(M_{gt}\):答案几何体
- \(\hat{A}_{i,t}\):优势估计,像就是好(正数),不像就是不好(负数)
- \(\epsilon\):裁剪阈值。这个与几何奖励提高训练稳定度有关。我们不希望模型进步与退步太快,类似于爬山算法。
- \(\beta\):惩罚权重,人为介入模型训练依据
- 为什么第二个GRPO损失公式不是“\(-E\dot\log()\)”形式:实际上,模型的进化与梯度有关。我们可以知道梯度与导数具有一致性,所以梯度运算符合求导法则,所以\(\log\)的增长率接近分数比例。我们可以用比例代替对数。事实上,对数是SFT,分数是裁剪强化学习。
- 过滤策略:把不符合几何特征、代码格式的代码过滤。这是在奖励函数之外独立存在的部分。
- 消融实验:把每一步删去后产生失败结果以证明该步骤的必要性。
0x01 SFT训练
0x01.1 配置环境
在modelscope notebook 中安装 ms-swift
# 安装/升级 ms-swift
pip install ms-swift -U
# 为了更好的训练效果和显存优化,建议安装以下可选依赖
pip install flash-attn --no-build-isolation # 加速注意力机制计算[reference:3]
pip install deepspeed # 支持多卡训练和ZeRO优化[reference:4]
0x01.2 准备训练数据
训练数据与LoRA训练数据区别较大。
本次复刻中采用了CoT思维链,重新构建了三维数据。实际上,现有的相关AI训练研究的数据集完全没有公开(公开的数据集大多不具有适配性或正确性保证),因此我们采取了生成常见几何体的方式进行训练。SFT训练数据(共300条)较为复杂,LoRA训练数据(共3000条)较为简单。
0x01.3 编写训练脚本
采用swift sft训练
# 设置使用的GPU,0代表第一张卡
CUDA_VISIBLE_DEVICES=0 \
swift sft \
--model Qwen/Qwen2.5-7B-Instruct \ # 指定基础模型[reference:8]
--train_type full \ # 【关键】指定为全参数微调[reference:9]
--dataset /path/to/your/train.jsonl \ # 数据集路径
--torch_dtype bfloat16 \ # 使用bfloat16精度,节省显存
--num_train_epochs 3 \ # 训练轮数
--per_device_train_batch_size 1 \ # 【关键】7B模型全参微调,batch_size设为1
--gradient_accumulation_steps 8 \ # 通过梯度累积来模拟更大的batch size
--learning_rate 2e-5 \ # 全参数微调常用学习率
--warmup_ratio 0.1 \ # 预热比例
--lr_scheduler_type cosine \ # 学习率调度器
--logging_steps 10 \ # 每10步打印一次日志
--save_steps 500 \ # 每500步保存一次检查点
--save_total_limit 2 \ # 只保留最近2个检查点,节省空间
--output_dir ./output_full_sft \ # 模型和日志的输出目录
--deepspeed default_zero2 \ # 【强烈推荐】使用DeepSpeed ZeRO-2优化显存[reference:10]
--max_length 2048 # 根据你的数据最大长度调整
0x01.4 自定义损失函数
通过ms-swift内置的loss_mapping字典注册,自主定义注册函数
# 在 swift/loss/mapping.py 文件中
def my_custom_loss(outputs, labels, loss_scale=None, num_items_in_batch=None):
# 在这里实现损失计算逻辑
# ...
return loss
loss_mapping = {
# ... 已有的损失函数映射 ...
"my_custom_loss": my_custom_loss, # 注册函数
}
在训练命令中,通过参数指令指定训练命令即可。
swift sft --loss_type my_custom_loss ...
训练开始后,观察日志中的loss值,如果loss持续下降,说明模型在学习。
0x01.5 训练结果展示



0x02 GRPO强化学习
本次训练分为两个阶段:
第一阶段stage1:使用TextToCad-2数据集,训练步数为120步,实现初步的学习代码编写技巧,保证代码可运行性。总耗时3小时20分钟。

第二阶段stage2:使用自定义数据集,实现对经典模型建模的学习(包括圆柱体(带中心孔)、矩形块(带圆角)、L 型支架、矩形板(带多个孔)、倒角、法兰盘等)。目前这部分遇到了比较严重的问题,还需要一段时间的训练。
0x02.1 奖励函数
0x02.1.1 stage1奖励函数
#!/usr/bin/env python3
"""
Stage 1 奖励函数: 简化版训练 (99% 时间)
目标: 快速收敛,让模型学会格式、语法、可执行性
奖励组成: 格式 + 语法 + 执行成功 + 代码相似度
"""
import re
import ast
import logging
from typing import Optional
from .cadquery_executor import get_executor
logger = logging.getLogger(__name__)
def extract_first_python_block(text: str) -> str:
"""提取第一个 python 代码块;没有则返回原文"""
match = re.search(r"```python\s*\n(.*?)```", text, re.DOTALL)
if match:
return match.group(1).strip()
return text.strip()
def clean_code_for_execution(code: str) -> str:
"""移除执行环境中不支持的显示函数调用"""
lines = code.split('\n')
cleaned = []
for line in lines:
stripped = line.strip()
if stripped.startswith('show_object(') or stripped.startswith('show('):
continue
cleaned.append(line)
return '\n'.join(cleaned)
def check_format(code: str) -> float:
score = 0.0
# 检查 thinking 块
if "```thinking" in code:
after_thinking = code.split("```thinking", 1)[1]
if "```" in after_thinking:
score += 0.25
# 检查 python 代码块
if "```python" in code:
score += 0.25
# 检查 CadQuery 导入
if "import cadquery as cq" in code or "import cadquery" in code:
score += 0.25
# 检查 result 赋值
if re.search(r"result\s*=", code):
score += 0.25
return score
# ============================================================
# 语法检查
# ============================================================
def check_syntax(code: str) -> float:
"""检查 Python 语法合法性"""
try:
py_code = extract_first_python_block(code)
ast.parse(py_code)
return 1.0
except SyntaxError:
return 0.0
except Exception:
return 0.0
# ============================================================
# 执行检查
# ============================================================
def check_executable(code: str) -> float:
"""
渐进式执行奖励:
0.0 语法/import 失败
0.1 运行时崩溃(至少 import 成功了)
0.2 执行完但 result 变量不存在
0.3 result 存在但 OCC 几何操作失败(API 用错了)
0.5 result 存在且是合法对象,但几何验证不通过(如空形体)
1.0 完全成功
"""
try:
py_code = extract_first_python_block(code)
py_code = clean_code_for_execution(py_code)
except Exception:
return 0.0
# 1. 语法检查
try:
ast.parse(py_code)
except SyntaxError:
return 0.0
# 2. 尝试执行
import os
try:
import cadquery as cq
import math
except Exception:
return 0.0
local_ns = {}
safe_globals = {
"__builtins__": __builtins__,
"cq": cq,
"math": math,
"os": os,
}
try:
exec(py_code, safe_globals, local_ns)
except Exception as e:
err = str(e)
print(f"[EXEC RUNTIME] {err[:200]}")
# 能跑到这里说明 import 和语法都没问题,只是 API 用错了
if "No module named" in err:
return 0.0
elif "must be planar" in err or "BRep_API" in err or "StdFail" in err:
return 0.30 # OCC 几何错误,给中等鼓励分
elif "has no attribute" in err or "'NoneType'" in err:
return 0.20 # 属性错误或空对象
else:
return 0.10 # 其他运行时错误,给基础分
# 3. 检查 result 变量
shape = local_ns.get("result")
if shape is None:
print("[EXEC] result variable not found")
return 0.20
# 4. 检查 result 类型和有效性
try:
if hasattr(shape, 'wrapped') and shape.wrapped is not None:
return 1.0
if hasattr(shape, 'Volume') and callable(shape.Volume) and shape.Volume() > 0:
return 1.0
if hasattr(shape, 'val') and shape.val() is not None:
return 1.0
# result 存在但无法验证为有效几何
print(f"[EXEC] result exists but invalid: type={type(shape)}")
return 0.50
except Exception as e:
print(f"[EXEC] validation error: {e}")
return 0.30
# ============================================================
# 代码相似度
# ============================================================
def code_similarity(generated: str, reference: str) -> float:
"""
计算生成代码与参考代码的相似度
使用 token 重叠率作为简单度量
"""
if not reference or not generated:
return 0.0
try:
gen_code = extract_first_python_block(generated)
ref_match = re.search(r"```python\s*\n(.*?)```", reference, re.DOTALL)
ref_code = ref_match.group(1) if ref_match else reference
gen_tokens = set(re.findall(r"\w+", gen_code.lower()))
ref_tokens = set(re.findall(r"\w+", ref_code.lower()))
if not ref_tokens:
return 0.0
overlap = len(gen_tokens & ref_tokens)
similarity = overlap / len(ref_tokens)
return min(similarity, 1.0)
except Exception:
return 0.0
# ============================================================
# 主奖励函数
# ============================================================
def compute_reward(code: str, prompt: str = "", stage: int = 1,
reference_code: str = "", target_mesh_path: str = "", **kwargs) -> float:
# 黑名单惩罚
bad_lib_keywords = ["pyvista", "pv.", "pytorch3d", "trimesh", "open3d",
"glut", "glbegin", "opengl", "cq_helper", "ch."]
penalty = 0.0
lower_code = code.lower()
for kw in bad_lib_keywords:
if kw in lower_code:
penalty -= 0.35
r_fmt = check_format(code)
r_syntax = check_syntax(code)
r_exec = check_executable(code)
r_sim = code_similarity(code, reference_code)
reward = (
0.15 * r_fmt +
0.15 * r_syntax +
0.60 * r_exec +
0.10 * r_sim
)
reward += penalty
reward = max(reward, 0.0)
# 精简版打印,只看关键
print(f"\n[REWARD] fmt={r_fmt:.2f} syntax={r_syntax:.2f} exec={r_exec:.2f} sim={r_sim:.2f} penalty={penalty:.2f} -> total={reward:.4f}\n")
return reward
0x02.1.2 stage2奖励函数
#!/usr/bin/env python3
"""
Stage 2 奖励函数: CD 几何奖励 + AST 代码结构相似度 fallback
"""
from __future__ import annotations # ✅ 关键:延迟求值类型注解,避免导入时解析 trimesh.Trimesh
import os
import re
import ast
import logging
import tempfile
from typing import Optional
import numpy as np
try:
import trimesh
from scipy.spatial import cKDTree
HAS_MESH_LIBS = True
except ImportError:
HAS_MESH_LIBS = False
trimesh = None # 占位,防止后面 type hint 报错
from .cadquery_executor import get_executor
from .cadquery_reward import extract_first_python_block, clean_code_for_execution
logger = logging.getLogger(__name__)
def check_format(code: str) -> float:
score = 0.0
if "```thinking" in code:
score += 0.25
if "```python" in code:
score += 0.25
if "import cadquery" in code:
score += 0.25
if re.search(r"result\s*[=]", code):
score += 0.25
return score
# ============================================================
# AST 代码结构相似度(当 mesh 不可用时 fallback)
# ============================================================
def ast_similarity(gen_code: str, ref_code: str) -> float:
try:
gen_tree = ast.parse(gen_code)
ref_tree = ast.parse(ref_code)
except SyntaxError:
return 0.0
def extract_features(node):
features = {'calls': [], 'methods': [], 'assigns': []}
for child in ast.walk(node):
if isinstance(child, ast.Call):
if isinstance(child.func, ast.Name):
features['calls'].append(child.func.id)
elif isinstance(child.func, ast.Attribute):
features['methods'].append(child.func.attr)
if isinstance(child, ast.Assign):
for target in child.targets:
if isinstance(target, ast.Name):
features['assigns'].append(target.id)
return features
gen_feat = extract_features(gen_tree)
ref_feat = extract_features(ref_tree)
gen_methods = set(gen_feat['methods'])
ref_methods = set(ref_feat['methods'])
method_sim = len(gen_methods & ref_methods) / len(ref_methods) if ref_methods else 0.0
gen_assigns = set(gen_feat['assigns'])
ref_assigns = set(ref_feat['assigns'])
assign_sim = len(gen_assigns & ref_assigns) / len(ref_assigns) if ref_assigns else 0.0
gen_calls = set(gen_feat['calls'])
ref_calls = set(ref_feat['calls'])
call_sim = len(gen_calls & ref_calls) / len(ref_calls) if ref_calls else 0.0
return min(0.5 * method_sim + 0.3 * assign_sim + 0.2 * call_sim, 1.0)
# ============================================================
# CD 相关函数
# ============================================================
def sample_point_cloud(mesh, n_points: int = 2048) -> np.ndarray:
if not HAS_MESH_LIBS:
raise RuntimeError("trimesh and scipy are required")
mesh_copy = mesh.copy()
vertices = np.array(mesh_copy.vertices)
center = (vertices.max(axis=0) + vertices.min(axis=0)) / 2.0
scale = np.linalg.norm(vertices.max(axis=0) - vertices.min(axis=0))
if scale < 1e-6:
scale = 1.0
normalized_vertices = (vertices - center) / scale
mesh_copy.vertices = normalized_vertices
points, _ = trimesh.sample.sample_surface(mesh_copy, n_points)
return points.astype(np.float32)
def chamfer_distance(P: np.ndarray, Q: np.ndarray) -> float:
if not HAS_MESH_LIBS:
raise RuntimeError("scipy is required")
tree_Q = cKDTree(Q)
dist_P_to_Q, _ = tree_Q.query(P, k=1)
term1 = np.mean(dist_P_to_Q ** 2)
tree_P = cKDTree(P)
dist_Q_to_P, _ = tree_P.query(Q, k=1)
term2 = np.mean(dist_Q_to_P ** 2)
return float(term1 + term2)
def cd_to_reward(cd: float) -> float:
if cd < 1e-5:
return 1.0
elif cd > 0.5:
return 0.0
else:
return max(0.0, 1.0 - cd * 1.98)
_target_mesh_cache = {}
def load_target_mesh(mesh_path: str) -> Optional[np.ndarray]:
if not mesh_path or not os.path.exists(mesh_path):
return None
if mesh_path in _target_mesh_cache:
return _target_mesh_cache[mesh_path]
try:
mesh = trimesh.load(mesh_path)
if isinstance(mesh, trimesh.Scene):
mesh = mesh.dump(concatenate=True)
points = sample_point_cloud(mesh, n_points=2048)
_target_mesh_cache[mesh_path] = points
return points
except Exception as e:
logger.warning(f"Failed to load target mesh {mesh_path}: {e}")
return None
def export_shape_to_mesh(shape):
"""导出 CadQuery shape 为 trimesh.Trimesh(无 mesh 库时返回 None)"""
if not HAS_MESH_LIBS:
return None
import cadquery as cq
with tempfile.NamedTemporaryFile(suffix=".stl", delete=False) as tmp:
tmp_path = tmp.name
try:
cq.exporters.export(shape, tmp_path)
pred_mesh = trimesh.load(tmp_path)
if isinstance(pred_mesh, trimesh.Scene):
pred_mesh = pred_mesh.dump(concatenate=True)
return pred_mesh
finally:
try:
os.unlink(tmp_path)
except:
pass
# ============================================================
# 主奖励函数
# ============================================================
def compute_reward(code: str, prompt: str = "", stage: int = 2,
reference_code: str = "", target_mesh_path: str = "", **kwargs) -> float:
py_code = extract_first_python_block(code)
py_code = clean_code_for_execution(py_code)
print(f"\n{'='*60}")
print(f"[EXTRACTED] first 300 chars:\n{py_code[:300]}")
print(f"[EXTRACTED] last 100 chars:\n{py_code[-100:]}")
print(f"[LENGTH] extracted={len(py_code)}")
r_fmt = 0.0
if "```python" in code or "import cadquery" in py_code:
r_fmt += 0.3
if "import cadquery" in py_code:
r_fmt += 0.3
if re.search(r"result\s*[=]", py_code):
r_fmt += 0.4
# 执行
success = False
shape = None
error = "Not executed"
try:
executor = get_executor(num_workers=1)
success, shape, error = executor.execute(py_code, timeout=30.0)
except Exception as e:
error = f"ExecutorException: {type(e).__name__}: {e}"
print(f"[EXEC EXCEPTION] {error}")
print(f"[EXEC] success={success}, error={error}")
# ========== 关键:部分奖励,让GRPO有信号 ==========
if success and shape is not None:
r_exec = 1.0
elif "SyntaxError" in error:
r_exec = 0.0 # 语法错误,0分
elif "Timeout" in error:
r_exec = 0.0 # 超时,0分
elif "Variable 'result' not found" in error:
r_exec = 0.1 # 忘了写result,给一点点
else:
# 运行时错误(API不存在、类型错误、ParseException等)
# 说明语法对了,逻辑/API错了,给鼓励让模型继续探索
r_exec = 0.3
print(f"[EXEC PARTIAL] Runtime error but syntax OK, giving 0.3")
# =====================================================
# CD 几何奖励(mesh 可用时)
r_geo = 0.0
has_mesh = False
if target_mesh_path and os.path.exists(target_mesh_path) and r_exec >= 1.0:
try:
pred_mesh = export_shape_to_mesh(shape)
if pred_mesh is None:
raise RuntimeError("Mesh libraries not available")
P = sample_point_cloud(pred_mesh, n_points=2048)
Q = load_target_mesh(target_mesh_path)
if Q is not None:
cd = chamfer_distance(P, Q)
r_geo = cd_to_reward(cd)
has_mesh = True
print(f"[STAGE2-CD] cd={cd:.6f} -> r_geo={r_geo:.4f}")
except Exception as e:
print(f"[STAGE2-CD-ERROR] {e}")
has_mesh = False
if not has_mesh:
r_sim = ast_similarity(py_code, reference_code)
print(f"[STAGE2-AST] sim={r_sim:.4f}")
# 语法OK但运行时错的给0.3 exec,让reward有区分度
reward = 0.05 * r_fmt + 0.55 * r_exec + 0.40 * r_sim
print(f"[STAGE2-REWARD] fmt={r_fmt:.2f} exec={r_exec:.2f} sim={r_sim:.4f} -> total={reward:.4f}")
return reward
reward = 0.05 * r_fmt + 0.15 * r_exec + 0.80 * r_geo
print(f"[STAGE2-REWARD] fmt={r_fmt:.2f} exec={r_exec:.2f} geo={r_geo:.4f} -> total={reward:.4f}")
return reward
0x02.2 训练基座train_GRPO
*这里的train_GRPO实际上是stage2的部分。stage1的GRPO比较普通,虽然留有备份但好像不太重要。
相较于cad-coder的训练模式,我们补充了训练的prompt以提高训练效率。实际上在训练过程中遇到了很多问题,例如ai输出过多陈述(思维链)占据了代码空间、训练过程中需要调整不同奖励函数的权重、拉取在线模型进行训练以减少训练成本等问题。由于本部分代码长度过长,这里只截取部分个人觉得重要的部分。
强制同步,直接在线拉取
# ========== 在 model = AutoModelForCausalLM.from_pretrained(...) 之后 ==========
# 强制同步 eos/pad token id,防止 QLoRA/4bit 加载后 config 丢失
if model.config.eos_token_id is None:
model.config.eos_token_id = tokenizer.eos_token_id
if model.config.pad_token_id is None:
model.config.pad_token_id = tokenizer.pad_token_id
# 对 Qwen 模型,有时 eos_token_id 会被设成 pad_token_id 的列表形式,强制修正
if isinstance(model.config.eos_token_id, list):
model.config.eos_token_id = model.config.eos_token_id[0]
if isinstance(model.config.pad_token_id, list):
model.config.pad_token_id = model.config.pad_token_id[0]
print(f"[MODEL CONFIG] eos_token_id={model.config.eos_token_id}, pad_token_id={model.config.pad_token_id}")
# ============================================================================
要求AI学会EOS停止符(事实上,这似乎是大多数LoRA训练法的通病)
# ========== 生成侧长度惩罚 + EOS偏置 ==========
generation_config = GenerationConfig(
max_new_tokens=grpo_args.max_completion_length,
do_sample=True,
temperature=grpo_args.temperature,
length_penalty=0.6,
)
generation_config.logit_bias = {tokenizer.eos_token_id: 7.0}
if "generation_config" in grpo_params:
grpo_kwargs["generation_config"] = generation_config
elif "generation_kwargs" in grpo_params:
grpo_kwargs["generation_kwargs"] = {
"length_penalty": 0.6,
"logit_bias": {tokenizer.eos_token_id: 2.0},
}
else:
if "length_penalty" in grpo_params:
grpo_kwargs["length_penalty"] = 0.6
logger.warning("当前TRL版本不支持完整生成配置,length_penalty已单独传入")
grpo_config = GRPOConfig(**grpo_kwargs)
logger.info(f"TRL GRPOConfig params used: {list(grpo_kwargs.keys())}")
prompt实践
prompt = (
"You are a CAD engineer. Given a text description, "
"plan the modeling steps and then write CadQuery code.\n\n"
f"Description: {text}\n\n"
"Please think step by step and then generate the code.\n"
"```thinking\n"
"1. Component decomposition\n"
"2. Coordinate system\n"
"3. Sketch design\n"
"4. Extrusion operations\n"
"5. Assembly\n"
"```\n\n"
"```python\n"
"import cadquery as cq\n"
)
0x02.3 训练脚本
*训练脚本分为两种,一种是初始训练用的脚本(train_stage1/2.sh),另一种是在训练中途中断后继续训练的脚本(resume_stage1/2.sh)。这里只展示继续训练的脚本。
这里的参数是需要反复调整的。主要的调整的参数是kl_coef(衡量参数之间差异)、num_generations(比较样本数)、max_completion_length(输出长度)、temperature(随机程度)。
0x02.3.1 stage1
#!/bin/bash
# ============================================================
# Stage 1: 从 checkpoint 恢复训练(防 hacking 参数版)
# ============================================================
set -e
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export PYTHONPATH="/mnt/workspace:${PYTHONPATH}"
# ===== 模型与路径配置 =====
SFT_BASE_MODEL="Qwen/Qwen2.5-3B-Instruct"
STAGE1_DIR="./output/stage_format"
REWARD_MODULE="plugin.cadquery_reward"
DATA_PATH="/mnt/workspace/final_train_sys.jsonl"
# ===== 自动查找最新的 checkpoint =====
LATEST_CKPT=$(ls -d ${STAGE1_DIR}/checkpoint-* 2>/dev/null | sort -V | tail -n 1)
if [ -z "${LATEST_CKPT}" ]; then
echo "Error: No checkpoint found in ${STAGE1_DIR}"
exit 1
fi
echo "========================================"
echo " Resuming Stage 1 from:"
echo " ${LATEST_CKPT}"
echo "========================================"
# ===== 启动训练(关键参数已调整防 hacking)=====
python train_grpo.py \
--stage 1 \
--model_name_or_path ${SFT_BASE_MODEL} \
--output_dir ${STAGE1_DIR} \
--reward_module ${REWARD_MODULE} \
--dataset_path ${DATA_PATH} \
--resume_from_checkpoint ${LATEST_CKPT} \
--max_steps 1500 \
--per_device_train_batch_size 1 \
--gradient_accumulation_steps 48 \
--learning_rate 2e-5 \
--kl_coef 0.01 \
--num_generations 6 \
--max_prompt_length 900 \
--max_completion_length 1024 \
--temperature 0.7 \
--load_in_4bit \
--use_qlora \
--lora_r 64 \
--lora_alpha 16 \
--bf16 \
--logging_steps 10 \
--save_steps 50 \
--save_total_limit 5 \
--eval_strategy no \
--report_to tensorboard \
--run_name "text2cad_stage1_resume" \
--seed 42
echo "Stage 1 resume complete!"
0x02.3.2 stage2
#!/bin/bash
# ============================================================
# Stage 2: 从 checkpoint 恢复训练
# ============================================================
set -e
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
# ===== 关键修复:添加项目根目录到 PYTHONPATH =====
export PYTHONPATH="/mnt/workspace:${PYTHONPATH}"
STAGE2_DIR="./output/stage2_grpo"
LATEST_CKPT=$(ls -d ${STAGE2_DIR}/checkpoint-* 2>/dev/null | sort -V | tail -n 1)
if [ -z "${LATEST_CKPT}" ]; then
echo "Error: No checkpoint found in ${STAGE2_DIR}"
exit 1
fi
echo "========================================"
echo " Resuming Stage 2 from:"
echo " ${LATEST_CKPT}"
echo "========================================"
REWARD_MODULE="plugin.cadquery_reward_paper"
DATA_PATH="/mnt/workspace/final_train_sys.jsonl"
python train_grpo.py \
--resume_from_checkpoint ${LATEST_CKPT} \
--model_name_or_path ${LATEST_CKPT} \
--output_dir ${STAGE2_DIR} \
--reward_module ${REWARD_MODULE} \
--dataset_path ${DATA_PATH} \
--max_steps 600 \
--per_device_train_batch_size 1 \
--gradient_accumulation_steps 48 \
--learning_rate 2e-6 \
--kl_coef 0.001 \
--num_generations 2 \
--max_prompt_length 768 \
--max_completion_length 1024 \
--temperature 0.1 \
--load_in_4bit \
--use_qlora \
--lora_r 64 \
--lora_alpha 16 \
--bf16 \
--logging_steps 10 \
--save_steps 200 \
--eval_strategy no \
--report_to tensorboard \
--run_name "text2cad_stage2_resume" \
--seed 42
echo "Stage 2 resume complete!"
0x03 模型封装
0x03.1 文件导出
我们计划在文件的最后添加一行代码:
cq.exporters.export(part, r"C:\Users\24961\Desktop\零件.step")
这样STEP即可直接导出到桌面,可以在FreeCAD直接打开。
0x03.2 文件封装
我们使用了图形化界面进行封装。设计了简洁的安装环境启动文件setup_env.bat与ai启动文件gui_run.bat。
实例:



浙公网安备 33010602011771号