lllyasviel / lllyasviel/ControlNet

self-trained ControlNet is slower than standard

Open
#683 0 comments 1 reaction 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
34.1k
Forks
3k
PR merge metrics
No merged PRs in 30d

Description

I train myself ControlNet according to [tutorial_train.py](https://github.com/lllyasviel/ControlNet/blob/main/tutorial_train.py).

After training, I got ```my_cn.ckpt```, size about 8G.

```my_cn.ckpt``` can load, run and get expected results by [gradio_scribble2image.py](https://github.com/lllyasviel/ControlNet/blob/main/gradio_scribble2image.py) , just update:
```python
model.load_state_dict(load_state_dict('./models/my_cn.ckpt', location='cuda'))
```

However, during inference, I found ```my_cn``` is ***several times slower*** than yours [```huggingface```](https://huggingface.co/lllyasviel/ControlNet/blob/main/models/control_sd15_scribble.pth).

I print ```state_dict``` in ```my_cn.ckpt``` and ```control_sd15_scribble.pth```, both are ```torch.float32```.

I test ControlNet alone, code as follow:
```python
from share import *
import cv2
import torch
from cldm.model import create_model, load_state_dict
from cldm.ddim_hacked import DDIMSampler
from tqdm import tqdm
from torchinfo import summary

model = create_model('./models/cldm_v15.yaml').cpu()

# 317 ms
# 1445.12M-params
sketch_ckpt_path = './models/control_sd15_scribble.pth'

# 1545 ms
# 1445.12M-params
# sketch_ckpt_path = './models/my_cn.ckpt'

model.load_state_dict(load_state_dict(sketch_ckpt_path, location='cuda'))

model = model.cuda()
control_net = model.control_model

x, hint, timesteps, context = torch.rand((1,4,64,64)).to('cuda'), torch.rand((1,3,512,512)).to('cuda'), torch.rand((1)).to('cuda'), torch.rand((1,77,768)).to('cuda')

# print model information: https://github.com/TylerYep/torchinfo
summary(control_net, input_data=[x, hint, timesteps, context])

epoch = 50
e_sum = 0.00
for i in tqdm(range(0, epoch)):
begin = cv2.getTickCount()
control_net(x, hint, timesteps, context)
end = cv2.getTickCount()
# to ms
e_sum += (end - begin) / cv2.getTickFrequency() * 1000.0
print(e_sum / epoch)
print("Done!")
```

I think I must have missed some details, looking forward to your suggestions.

Contributor guide

No contributing guide indexed for this repository

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start by reproducing the timing comparison in the issue with control_sd15_scribble.pth and my_cn.ckpt, then inspect tutorial_train.py, gradio_scribble2image.py, and the model-loading path used by cldm.model. Done means identifying the cause of the checkpoint-dependent slowdown and documenting a verified fix or explanation.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, performance
Issue type
Bug
Difficulty
4/5
Estimated time
3-5 days
Activity status
Stale
Clarity
Needs clarification
Newbie friendliness
25/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.