RAFT代码

论坛 期权论坛     
选择匿名的用户   2021-5-22 16:19   37   0
<p>corr 里面的torch.matmul</p>
<h1>总的代码</h1>
<pre class="blockcode"><code class="language-python">def train(args):

    model &#61; nn.DataParallel(RAFT(args), device_ids&#61;args.gpus)
    print(&#34;Parameter Count: %d&#34; % count_parameters(model))

    if args.restore_ckpt is not None:
        model.load_state_dict(torch.load(args.restore_ckpt), strict&#61;False)

    model.cuda()
    model.train()

    if args.stage !&#61; &#39;chairs&#39;:
        model.module.freeze_bn()

    train_loader &#61; datasets.fetch_dataloader(args)
    optimizer, scheduler &#61; fetch_optimizer(args, model)

    total_steps &#61; 0
    scaler &#61; GradScaler(enabled&#61;args.mixed_precision)
    logger &#61; Logger(model, scheduler)

    VAL_FREQ &#61; 5000
    add_noise &#61; True

    should_keep_training &#61; True
    while should_keep_training:

        for i_batch, data_blob in enumerate(train_loader):
            optimizer.zero_grad()
            image1, image2, flow, valid &#61; [x.cuda() for x in data_blob]

            if args.add_noise:
                stdv &#61; np.random.uniform(0.0, 5.0)
                image1 &#61; (image1 &#43; stdv * torch.randn(*image1.shape).cuda()).clamp(0.0, 255.0)
                image2 &#61; (image2 &#43; stdv * torch.randn(*image2.shape).cuda()).clamp(0.0, 255.0)

            flow_predictions &#61; model(image1, image2, iters&#61;args.iters)            

            loss, metrics &#61; sequence_loss(flow_predictions, flow, valid, args.gamma)
            scaler.scale(loss).backward()
            scaler.unscale_(optimizer)               
            torch.nn.utils.clip_grad_norm_(model.parameters(), args.clip)
            
            scaler.step(optimizer)
            scheduler.step()
            scaler.update()

            logger.push(metrics)

            if total_steps % VAL_FREQ &#61;&#61; VAL_FREQ - 1:
                PATH &#61; &#39;checkpoints/%d_%s.pth&#39; % (total_steps&#43;1, args.name)
                torch.save(model.state_dict(), PATH)

                results &#61; {}
                for val_dataset in args.validation:
                    if val_dataset &#61;&#61; &#39;chairs&#39;:
                        results.update(evaluate.validate_chairs(model.module))
                    elif val_dataset &#61;&#61; &#39;sintel&#39;:
                        results.update(evaluate.validate_sintel(model.module))
                    elif val_dataset &#61;&#61; &#39;kitti&#39;:
                        results.update(evaluate.validate_kitti(model.module))

                logger.write_dict(results)
               
                model.train()
                if args.stage !&#61; &#39;chairs&#39;:
                    model.module.freeze_bn()
            
            total_steps &#43;&#61; 1

            if total_steps &gt; args.num_steps:
                should_keep_training &#61; False
                break

    logger.close()
    PATH &#61; &#39;checkpoints/%s.pth&#39; % args.name
    torch.save(model.state_dict(), PATH)

    return PATH</code></pre>
<h2>首先进行网络初始化</h2>
<p>首先进行RAFT的初始化:有一个选项为args.small。</p>
<pre class="blockcode"><code class="language-python">class RAFT(nn.Module):
    def __init__(self, args):
        super(RAFT, self).__init__()
        self.args &#61; args

        if args.small:
            self.hidden_dim &#61; hdim &#61; 96
            self.context_dim &#61; cdim &#61; 64
            args.corr_levels &#61; 4
            args.corr_radius &#61; 3
        
        else:
            self.hidden_dim &#61; hdim &#61; 128
            self.context_dim &#61; cdim &#61; 128
            args.corr_levels &#61; 4
            args.corr_radius &#61; 4</code></pre>
<p>然后进行网络的初始化</p>
<pre class="blockcode"><code class="language-python">        if args.small:
            self.fnet &#61; SmallEncoder(output_dim&#61;128, norm_fn&#61;&#39;instance&#39;, dropout&#61;args.dropout)        
            self.cnet &#61; SmallEncoder(output_dim&#61;hdim&#43;cdim, norm_fn&#61;&#39;none&#39;, dropout&#61;args.dropout)
            self.update_block &#61; SmallUpdateBlock(self.args, hidden_dim&#61;hdim)

        else:
            self.fnet &#61; BasicEncoder(output_dim&#61;256, norm_fn&#61;&#39;instance&#39;, dropout&#61;args.dropout)        
            self.cnet &#61; BasicEncoder(output_dim&#61;hdim&#43;cdim, norm_fn&#61;&#39;batch&#39;, dropout&#61;args.dropout)
            self.update_block &#61; BasicUpdateBlock(self.args, hidden_dim&#61;hdim)</code></pre>
<p>然后开始进行basicEncoder的初始化,默认instance</p>
<pre class="blockcode"><code class="language-python">class BasicEncoder(nn.Module):
    def __init__(self, output_dim&#61;128, norm_fn&#61;&#39;batch&#39;, dropout&#61;0.0):
        super(BasicEncoder, self).__init__()
        self.norm_fn &#61; norm_fn

        if self.norm_fn &#61;&#61; &#39;group&#39;:
            self.norm1 &#61; nn.GroupNorm(num_groups&#61;8, num_channels&#61;64)
            
        elif self.norm_fn &#61;&#61; &#39;batch&#39;:
            self.norm1 &#61; nn.BatchNorm2d(64)

        elif self.norm_fn &#61;&#61; &#39;instance&#39;:
            self.norm1 &#61; nn.InstanceNorm2d(64)

        self.conv1 &#61; nn.Conv2d(3, 64, kernel_size&#61;7, stride&#61;2, padding&#61;3)
        self.relu1 &#61; nn.ReLU(inplace&#61;True)

        self.in_planes &#61; 64
        self.layer1 &#61; self._make_layer(64,  stride&#61;1)
        self.layer2 &#61; self._make_layer(96, stride&#61;2)
        self.layer3 &#61; self._make_layer(128, stride&#61;2)

          # output convol
分享到 :
0 人收藏
您需要登录后才可以回帖 登录 | 立即注册

本版积分规则

积分:3875789
帖子:775174
精华:0
期权论坛 期权论坛
发布
内容

下载期权论坛手机APP