text2CAD模型复刻实践记录

0x00 论文概述

0x00.1 生成流程

由冷启动初步训练SFT模型,随后从几何体角度使用CoT强化学习训练模型。强化学习涉及奖励函数(相对优势\(\hat{A}\))与数据集构建部分,奖励函数由倒角距离CD,第二个损失函数有关。

0x00.2 主要专业术语

  1. GRPO:群体奖励策略优化(一个模型得到多个答案选最优)
  2. 几何奖励\(R_i^{geo}\):模型生成的几何体与答案几何体是否一致(由倒角距离CD衡量)
  3. 格式奖励\(R_i^{fmt}\):生成代码语法是否正确
  4. CoT:思维链(从策略规划到代码生成)
  5. SFT:监督微调(照着答案改)
  6. 命令序列S:模型模仿人类构建模型的思考过程
  7. L:自然语言提示词
  8. \(C_{gt}\) :正确Cadquery代码(一般以token方式处理)
  9. 损失函数:由于模型生成产生的与答案几何体的偏差。一般损失越小,答案越优。
  10. \(E_{(L,C_{gt})\sim D}\):期望
  11. 冷启动:第一轮学习前对模型进行优质数据的监督微调
  12. 倒角距离CD:双向找最近点。对于预测点阵与答案点阵中的每一个点,都找到另一个点阵相同坐标附近的最近点,把所有距离求和取平均(或者别的方法)
  13. \(||x-y||^2_2\):点x与点y之间的欧几里得距离
  14. \(P\):预测点阵
  15. \(Q\):答案点阵
  16. \(L_{cot}\):规划流程
  17. \(M_{gt}\):答案几何体
  18. \(\hat{A}_{i,t}\):优势估计,像就是好(正数),不像就是不好(负数)
  19. \(\epsilon\):裁剪阈值。这个与几何奖励提高训练稳定度有关。我们不希望模型进步与退步太快,类似于爬山算法。
  20. \(\beta\):惩罚权重,人为介入模型训练依据
  21. 为什么第二个GRPO损失公式不是“\(-E\dot\log()\)”形式:实际上,模型的进化与梯度有关。我们可以知道梯度与导数具有一致性,所以梯度运算符合求导法则,所以\(\log\)的增长率接近分数比例。我们可以用比例代替对数。事实上,对数是SFT,分数是裁剪强化学习。
  22. 过滤策略:把不符合几何特征、代码格式的代码过滤。这是在奖励函数之外独立存在的部分。
  23. 消融实验:把每一步删去后产生失败结果以证明该步骤的必要性。

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 训练结果展示

image
loss_curve
image

0x02 GRPO强化学习

本次训练分为两个阶段:
第一阶段stage1:使用TextToCad-2数据集,训练步数为120步,实现初步的学习代码编写技巧,保证代码可运行性。总耗时3小时20分钟。
image
第二阶段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。
实例:

image

image

posted @ 2026-07-21 18:20  adolf_stalin  阅读(40)  评论(1)    收藏  举报