mirror of
https://github.com/wassname/bert_adapter.git
synced 2026-08-03 12:40:53 +08:00
Changes in configuration.
This commit is contained in:
@@ -36,9 +36,9 @@ class TokenizedDataFrameDataset(Dataset):
|
||||
def preprocess_text(self, text: str) -> BertInput:
|
||||
tokens = self.tokenizer.tokenize(text)
|
||||
tokens = ["[CLS]"] + tokens + ["[SEP]"]
|
||||
input_ids = self.tokenizer.convert_tokens_to_ids(tokens)
|
||||
input_ids = self.tokenizer.convert_tokens_to_ids(tokens)[:self.max_seq_len]
|
||||
|
||||
segment_ids = [0] * len(tokens)
|
||||
segment_ids = [0] * len(input_ids)
|
||||
input_mask = [1] * len(input_ids)
|
||||
|
||||
padding = [0] * (self.max_seq_len - len(input_ids))
|
||||
@@ -46,16 +46,15 @@ class TokenizedDataFrameDataset(Dataset):
|
||||
input_mask += padding
|
||||
segment_ids += padding
|
||||
|
||||
assert len(input_ids) == self.max_seq_len
|
||||
assert len(input_mask) == self.max_seq_len
|
||||
assert len(segment_ids) == self.max_seq_len
|
||||
|
||||
assert len(input_ids) == self.max_seq_len, f'{len(input_ids)} != {self.max_seq_len}'
|
||||
assert len(input_mask) == self.max_seq_len, f'{len(input_mask)} != {self.max_seq_len}'
|
||||
assert len(segment_ids) == self.max_seq_len, f'{len(segment_ids)} != {self.max_seq_len}'
|
||||
return BertInput(*[torch.LongTensor(_) for _ in [input_ids, input_mask, segment_ids]])
|
||||
|
||||
def preprocess_label(self, label: int):
|
||||
result = torch.zeros(len(self.y_labels))
|
||||
result[label] = 1
|
||||
return result
|
||||
return torch.LongTensor(result)
|
||||
|
||||
def __getitem__(self, index) -> dict:
|
||||
sample = self.df.iloc[index]
|
||||
|
||||
+1
-1
@@ -92,7 +92,7 @@ class ClassificationModel(nn.Module):
|
||||
|
||||
if labels is not None:
|
||||
loss_function = nn.CrossEntropyLoss()
|
||||
loss = loss_function(logits.view(-1, self.n_labels), labels.view(-1))
|
||||
loss = loss_function(logits, labels)
|
||||
return loss, logits
|
||||
else:
|
||||
return logits
|
||||
|
||||
+7
-2
@@ -64,7 +64,7 @@ if __name__ == '__main__':
|
||||
model.eval()
|
||||
model.to(device)
|
||||
|
||||
optimizer = Adam(model.learnable_parameters(), lr=0.001, amsgrad=True)
|
||||
optimizer = Adam(model.parameters(), lr=0.001, amsgrad=True)
|
||||
|
||||
print('Model have initialized')
|
||||
for i in range(args.num_epochs):
|
||||
@@ -73,7 +73,12 @@ if __name__ == '__main__':
|
||||
optimizer.zero_grad()
|
||||
|
||||
input_ids, input_mask, segment_ids = batch['x']
|
||||
input_ids = input_ids.to(device)
|
||||
input_mask = input_mask.to(device)
|
||||
segment_ids = segment_ids.to(device)
|
||||
|
||||
y = batch['y']
|
||||
y = y.to(device)
|
||||
|
||||
loss, _ = model.forward(input_ids, input_mask, segment_ids, labels=y)
|
||||
loss.backward()
|
||||
@@ -83,7 +88,7 @@ if __name__ == '__main__':
|
||||
model.eval()
|
||||
labels = []
|
||||
predictions = []
|
||||
for batch in test_loader:
|
||||
for batch in tqdm(test_loader):
|
||||
optimizer.zero_grad()
|
||||
|
||||
input_ids, input_mask, segment_ids = batch['x']
|
||||
|
||||
Reference in New Issue
Block a user