# VirtualTryOn **Repository Path**: WindNegev/virtual-try-on ## Basic Information - **Project Name**: VirtualTryOn - **Description**: 基于stable diffusion的虚拟试穿 - **Primary Language**: Python - **License**: Not specified - **Default Branch**: master - **Homepage**: None - **GVP Project**: No ## Statistics - **Stars**: 0 - **Forks**: 1 - **Created**: 2025-12-13 - **Last Updated**: 2026-07-15 ## Categories & Tags **Categories**: Uncategorized **Tags**: None ## README # Virtual Try-On 项目 基于 CatVTON 的虚拟试衣项目,支持服装试穿功能的训练和推理。 ## 项目概述 本项目实现了基于扩散模型的虚拟试衣系统,支持: - 服装试穿训练(DreamBooth + Inpainting) - CatVTON 架构适配 - 多种条件控制(OpenPose、Canny边缘检测) - 变形服装(Warp Cloth)支持 ## 环境要求 - Python 3.8+ - PyTorch 1.10+ - CUDA 11.0+(推荐使用GPU) - 8GB+ GPU显存(推荐16GB+) ## 安装步骤 ### 1. 克隆项目 ```bash git clone https://github.com/yourusername/virtual-try-on.git cd virtual-try-on ``` ### 2. 安装依赖 ```bash pip install -r requirements.txt ``` 主要依赖包包括: - diffusers - transformers - accelerate - torch - torchvision - pillow - numpy - tqdm - safetensors ### 3. 下载预训练模型 #### 方法一:通过 Hugging Face Hub(推荐) ```bash # 安装 huggingface_hub pip install huggingface_hub # 下载 CatVTON 模型 python -c " from huggingface_hub import snapshot_download snapshot_download(repo_id='zhengchong/CatVTON', local_dir='./models/catvton') " ``` #### 方法二:手动下载 1. 访问 [CatVTON Hugging Face 页面](https://huggingface.co/zhengchong/CatVTON) 2. 下载以下文件: - `model.safetensors` - CatVTON注意力权重 - `config.json` - 模型配置 3. 将下载的文件放置到 `./models/catvton/` 目录 #### 方法三:使用镜像(国内用户) ```bash # 设置 Hugging Face 镜像 export HF_ENDPOINT="https://hf-mirror.com" # 然后使用方法一下载 ``` ### 4. 下载基础 Stable Diffusion 模型 ```bash # 下载 Stable Diffusion 1.5 python -c " from huggingface_hub import snapshot_download snapshot_download(repo_id='runwayml/stable-diffusion-v1-5', local_dir='./models/stable-diffusion-v1-5') " # 下载 VAE 模型 python -c " from huggingface_hub import snapshot_download snapshot_download(repo_id='stabilityai/sd-vae-ft-mse', local_dir='./models/sd-vae-ft-mse') " ``` ## 项目结构 ``` virtual-try-on/ ├── train_dreambooth_inpaint_catvton_base.py # 主训练脚本 ├── catvton_base_infer.py # 推理脚本 ├── unet_adapter.py # UNet适配器 ├── models/ # 模型目录 │ ├── catvton/ # CatVTON模型 │ ├── stable-diffusion-v1-5/ # Stable Diffusion基础模型 │ └── sd-vae-ft-mse/ # VAE模型 ├── data/ # 数据目录 │ ├── real_images/ # 真人图像 │ ├── real_masks/ # 区域掩码 │ ├── condition_images/ # 条件图像(服装) │ ├── cloth_warp_images/ # 变形服装图像(可选) │ ├── cloth_warp_masks/ # 变形服装掩码(可选) │ ├── openpose_images/ # OpenPose姿态图像(可选) │ └── canny_images/ # Canny边缘图像(可选) └── outputs/ # 输出目录 ``` ## 数据准备 ### 训练数据格式 训练数据需要按以下结构组织: ``` your_dataset/ ├── real_images/ # 真人图像 │ ├── 001.jpg │ ├── 002.jpg │ └── ... ├── real_masks/ # 试穿区域掩码(黑白图像) │ ├── 001.png │ ├── 002.png │ └── ... ├── condition_images/ # 条件图像(服装图像) │ ├── 001.jpg │ ├── 002.jpg │ └── ... └── [可选目录] ├── cloth_warp_images/ # 变形后的服装图像 ├── cloth_warp_masks/ # 变形服装掩码 ├── openpose_images/ # OpenPose姿态图像 └── canny_images/ # Canny边缘检测图像 ``` **注意:** - 所有文件夹中的图像数量必须相同 - 图像文件名需要一一对应 - 掩码图像为黑白图像,白色区域表示需要替换的区域 - 图像分辨率建议为 512x512 ## 使用方法 ### 1. 训练模型 ```bash python train_dreambooth_inpaint_catvton_base.py \ --pretrained_model_name_or_path="./models/stable-diffusion-v1-5" \ --catvton_attn_path="./models/catvton" \ --instance_data_dir="./data/your_dataset" \ --instance_prompt="a person wearing clothes" \ --output_dir="./outputs" \ --resolution=512 \ --train_batch_size=4 \ --gradient_accumulation_steps=1 \ --gradient_checkpointing \ --learning_rate=5e-6 \ --max_train_steps=5000 \ --validation_steps=500 \ --validation_prompt="a person wearing fashionable clothes" \ --validation_image="./data/validation/real_image.jpg" \ --validation_mask="./data/validation/mask.png" \ --validation_condition_image="./data/validation/cloth.jpg" \ --use_warp_cloth \ --use_openpose_conditioning \ --use_canny_conditioning \ --latent_append_num=3 \ --save_steps=1000 \ --save_total_limit=3 ``` #### 主要训练参数说明 - `--pretrained_model_name_or_path`: Stable Diffusion基础模型路径 - `--catvton_attn_path`: CatVTON注意力权重路径 - `--instance_data_dir`: 训练数据目录 - `--instance_prompt`: 实例提示词 - `--output_dir`: 输出目录 - `--resolution`: 图像分辨率(默认512) - `--train_batch_size`: 训练批次大小 - `--learning_rate`: 学习率 - `--max_train_steps`: 最大训练步数 - `--use_warp_cloth`: 是否使用变形服装 - `--use_openpose_conditioning`: 是否使用OpenPose条件 - `--use_canny_conditioning`: 是否使用Canny边缘条件 - `--latent_append_num`: 使用的latent数量(1=CatVTON, 2=CatVTON+Openpose, 3=CatVTON+Canny+Openpose) ### 2. 模型推理 ```python from catvton_base_infer import run_inference from diffusers import AutoencoderKL, DDPMScheduler, UNet2DConditionModel import torch # 加载模型 unet = UNet2DConditionModel.from_pretrained("./models/stable-diffusion-v1-5", subfolder="unet") vae = AutoencoderKL.from_pretrained("./models/sd-vae-ft-mse") noise_scheduler = DDPMScheduler.from_pretrained("./models/stable-diffusion-v1-5", subfolder="scheduler") # 加载训练好的权重 from unet_adapter import adapt_unet_with_catvton_attn adapt_unet_with_catvton_attn(unet, "./outputs/your_checkpoint") # 运行推理 result = run_inference( unet=unet, vae=vae, noise_scheduler=noise_scheduler, device='cuda', image="path/to/person_image.jpg", mask="path/to/mask.png", condition_image="path/to/cloth_image.jpg", num_inference_steps=50, guidance_scale=2.5 )[0] result.save("output.jpg") ``` ### 3. 恢复训练 ```bash python train_dreambooth_inpaint_catvton_base.py \ # ... 其他参数 ... --resume_from_checkpoint="./outputs/step-2000" \ # 或者使用最新检查点 # --resume_from_checkpoint="latest" ``` ## 性能优化建议 1. **内存优化**: - 使用 `--gradient_checkpointing` 减少显存使用 - 使用 `--use_8bit_adam` 优化器 - 减小 `--train_batch_size` 并增加 `--gradient_accumulation_steps` 2. **训练加速**: - 使用 `--mixed_precision fp16` 启用混合精度训练 - 使用多GPU训练:`accelerate launch --multi_gpu` 3. **网络优化**(国内用户): - 设置 HF_ENDPOINT 环境变量使用镜像 - 或者手动下载模型文件 ## 常见问题 ### Q: 显存不足怎么办? A: - 使用 `--gradient_checkpointing` - 减小批次大小 `--train_batch_size` - 使用 `--mixed_precision fp16` - 使用 `--use_8bit_adam` ### Q: 训练速度很慢? A: - 确保使用GPU训练 - 启用混合精度训练 - 考虑使用多GPU训练 ### Q: 模型效果不好? A: - 增加训练步数 - 调整学习率 - 使用更高质量的训练数据 - 尝试不同的条件控制组合 ## 许可证 本项目基于 MIT 许可证开源。 ## 致谢 - [CatVTON](https://huggingface.co/zhengchong/CatVTON) - 提供虚拟试衣的基础架构 - [Diffusers](https://github.com/huggingface/diffusers) - 扩散模型实现 - [Stable Diffusion](https://huggingface.co/runwayml/stable-diffusion-v1-5) - 基础生成模型