TypeError: forward() takes 3 positional arguments but 5 were given when add GradScaler() amp hook

Viewed 520

I am training resnet34 for handwriting recognition and want to optimize the training time of the model using GradScaler (). I initialize AMP in the train_loop function which is responsible for starting the train session. My output should be enc_pad_texts, output_lenghts, text_lens but feeding them to loss_fp gives me an error in foward. How can I rewrite foward or change loss_fn

a function that helps to combine images and target text into a batch
def collate_fn(batch):
    images, texts, enc_texts = zip(*batch)
    images = torch.stack(images, 0)
    text_lens = torch.LongTensor([len(text) for text in texts])
    enc_pad_texts = pad_sequence(enc_texts, batch_first=True, padding_value=0)
    return images, texts, enc_pad_texts, text_lens
from torch.cuda.amp import GradScaler
from torch.cuda.amp import autocast

loss_fn = nn.CrossEntropyLoss()

def train_loop(data_loader, model, criterion, optimizer, epoch):

    torch.autograd.set_detect_anomaly(False)
    torch.autograd.profiler.profile(False)
    torch.autograd.profiler.emit_nvtx(False)

    scaler = GradScaler()

    torch.backends.cudnn.benchmark = True
    loss_avg = AverageMeter()
    model.train()

    for images, texts, enc_pad_texts, text_lens in data_loader:

        model.zero_grad()
        images = images.to(DEVICE)
        batch_size = len(texts)
        with autocast():
          output = model(images)

          output_lenghts = torch.full(
              size=(output.size(1),),
              fill_value=output.size(0),
              dtype=torch.long
          )

          loss = criterion(output, enc_pad_texts, output_lenghts, text_lens)
          loss_avg.update(loss.item(), batch_size)

          loss = loss_fn(output, enc_pad_texts, output_lenghts, text_lens)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    
    torch.nn.utils.clip_grad_norm_(model.parameters(), 2)
    scaler.update()

    for param_group in optimizer.param_groups:
        lr = param_group['lr']
    print(f'\nEpoch {epoch}, Loss: {loss_avg.avg:.5f}, LR: {lr:.7f}')
    return loss_avg.avg
def get_resnet34_backbone(pretrained=True):
    m = torchvision.models.resnet34(pretrained=True)
    input_conv = nn.Conv2d(3, 64, 7, 1, 3)
    blocks = [input_conv, m.bn1, m.relu,
              m.maxpool, m.layer1, m.layer2, m.layer3]
    return nn.Sequential(*blocks)


class BiLSTM(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, dropout=0.1):
        super().__init__()
        self.lstm = nn.LSTM(
            input_size, hidden_size, num_layers,
            dropout=dropout, batch_first=True, bidirectional=True)

    def forward(self, x):
        out, _ = self.lstm(x)
        return out


class CRNN(nn.Module):
    def __init__(
        self, number_class_symbols, time_feature_count=256, lstm_hidden=256,
        lstm_len=2,
    ):
        super().__init__()
        self.feature_extractor = get_resnet34_backbone(pretrained=True)
        self.avg_pool = nn.AdaptiveAvgPool2d(
            (time_feature_count, time_feature_count))
        self.bilstm = BiLSTM(time_feature_count, lstm_hidden, lstm_len)
        self.classifier = nn.Sequential(
            nn.Linear(lstm_hidden * 2, time_feature_count),
            nn.GELU(),
            nn.Dropout(0.01),
            nn.Linear(time_feature_count, number_class_symbols)
        )

    def forward(self, x):
        x = self.feature_extractor(x)
        b, c, h, w = x.size()
        x = x.view(b, c * h, w)
        x = self.avg_pool(x)
        x = x.transpose(1, 2)
        x = self.bilstm(x)
        x = self.classifier(x)
        x = nn.functional.log_softmax(x, dim=2).permute(1, 0, 2)
        return x
<ipython-input-39-09956edc294f> in train(config)
    122     for epoch in range(config['num_epochs']):
    123 
--> 124         loss_avg = train_loop(train_loader, model, criterion, optimizer, epoch)
    125         acc_avg = val_loop(val_loader, model, tokenizer, DEVICE)
    126         scheduler.step(acc_avg)

<ipython-input-39-09956edc294f> in train_loop(data_loader, model, criterion, optimizer, epoch)
     42           loss_avg.update(loss.item(), batch_size)
     43 
---> 44           loss = loss_fn(output,  enc_pad_texts, output_lenghts, text_lens)
     45 
     46     scaler.scale(loss).backward()

/usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs)
   1049         if not (self._backward_hooks or self._forward_hooks or self._forward_pre_hooks or _global_backward_hooks
   1050                 or _global_forward_hooks or _global_forward_pre_hooks):
-> 1051             return forward_call(*input, **kwargs)
   1052         # Do not call functions when jit is used
   1053         full_backward_hooks, non_full_backward_hooks = [], []

TypeError: forward() takes 3 positional arguments but 5 were given

Train func


def train(config):
    
    tokenizer = Tokenizer(config['alphabet'])
    os.makedirs(config['save_dir'], exist_ok=True)
    train_loader, val_loader = get_loaders(tokenizer, config)

    model = CRNN(number_class_symbols=tokenizer.get_num_chars())
    model.load_state_dict(torch.load("/content/drive/MyDrive/model-1-0.6960.ckpt"))
    model.to(DEVICE)

    criterion = torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=True)
    optimizer = torch.optim.AdamW(model.parameters(), lr=0.0001,
                                  weight_decay=0.1)
    scheduler = scheduler = torch.optim.lr_scheduler.OneCycleLR(
        optimizer=optimizer,
        epochs=config.get('num_epochs'),
        steps_per_epoch=len(train_loader),
        max_lr=0.01,
        pct_start=0.07,
        anneal_strategy='cos',
        final_div_factor=10 ** 3
    )

    best_acc = -np.inf
    acc_avg = val_loop(val_loader, model, tokenizer, DEVICE)
    
    for epoch in range(config['num_epochs']):

        loss_avg = train_loop(train_loader, model, criterion, optimizer, epoch)
        acc_avg = val_loop(val_loader, model, tokenizer, DEVICE)
        scheduler.step(acc_avg)
        print(f"Epoch: {epoch} Loss_avg: {loss_avg} Acc_avg: {acc_avg} Step {scheduler.step(acc_avg)}" )
        if acc_avg > best_acc:
            best_acc = acc_avg
            model_save_path = os.path.join(
                config['save_dir'], f'model-{epoch}-{acc_avg:.4f}.ckpt')
            torch.save(model.state_dict(), model_save_path)
            print('Model weights saved')
0 Answers
Related