Compare commits

...
7 Commits
2 changed files with 53 additions and 26 deletions
+6
View File
@@ -24,6 +24,7 @@ import os
import logging
import argparse
import random
from datetime import datetime
from tqdm import tqdm, trange
import numpy as np
@@ -580,6 +581,7 @@ def main():
model.eval()
eval_loss, eval_accuracy = 0, 0
nb_eval_steps, nb_eval_examples = 0, 0
start_time = datetime.now()
for input_ids, input_mask, segment_ids, label_ids in eval_dataloader:
input_ids = input_ids.to(device)
input_mask = input_mask.to(device)
@@ -600,6 +602,10 @@ def main():
nb_eval_examples += input_ids.size(0)
nb_eval_steps += 1
total_time = datetime.now() - start_time
logger.info("Total time: %d seconds — Number of examples: %d — Examples/second: %f" % (
total_time.total_seconds(), len(eval_data), len(eval_data)/total_time.total_seconds()))
eval_loss = eval_loss / nb_eval_steps
eval_accuracy = eval_accuracy / nb_eval_examples
+47 -26
View File
@@ -27,6 +27,7 @@ import math
import os
import random
import pickle
from datetime import datetime
from tqdm import tqdm, trange
import numpy as np
@@ -683,9 +684,11 @@ def main():
help="Bert pre-trained model selected in the list: bert-base-uncased, "
"bert-large-uncased, bert-base-cased, bert-base-multilingual, bert-base-chinese.")
parser.add_argument("--output_dir", default=None, type=str, required=True,
help="The output directory where the model checkpoints and predictions will be written.")
help="The output directory where the predictions will be written.")
## Other parameters
parser.add_argument("--model_save_dir", default=None, type=str,
help="The directory where the model checkpoints will be saved and loaded from.")
parser.add_argument("--train_file", default=None, type=str, help="SQuAD json for training. E.g., train-v1.1.json")
parser.add_argument("--predict_file", default=None, type=str,
help="SQuAD json for predictions. E.g., dev-v1.1.json or test-v1.1.json")
@@ -699,6 +702,7 @@ def main():
"be truncated to this length.")
parser.add_argument("--do_train", default=False, action='store_true', help="Whether to run training.")
parser.add_argument("--do_predict", default=False, action='store_true', help="Whether to run eval on the dev set.")
parser.add_argument("--do_benchmark", default=False, action='store_true', help="Whether to benchmark prediction speed on the dev set.")
parser.add_argument("--train_batch_size", default=32, type=int, help="Total batch size for training.")
parser.add_argument("--predict_batch_size", default=8, type=int, help="Total batch size for predictions.")
parser.add_argument("--learning_rate", default=5e-5, type=float, help="The initial learning rate for Adam.")
@@ -740,6 +744,10 @@ def main():
default=False,
action='store_true',
help="Whether to use 16-bit float precision instead of 32-bit")
parser.add_argument('--predict_fp16',
default=False,
action='store_true',
help="Whether to use 16-bit float precision instead of 32-bit for prediction")
parser.add_argument('--loss_scale',
type=float, default=0,
help="Loss scaling to improve fp16 numeric stability. Only used when fp16 set to True.\n"
@@ -787,6 +795,9 @@ def main():
if os.path.exists(args.output_dir) and os.listdir(args.output_dir):
raise ValueError("Output directory () already exists and is not empty.")
os.makedirs(args.output_dir, exist_ok=True)
if not args.model_save_dir:
args.model_save_dir = args.output_dir
output_model_file = os.path.join(args.model_save_dir, "pytorch_model.bin")
tokenizer = BertTokenizer.from_pretrained(args.bert_model)
@@ -915,17 +926,18 @@ def main():
optimizer.zero_grad()
global_step += 1
# Save a trained model
model_to_save = model.module if hasattr(model, 'module') else model # Only save the model it-self
output_model_file = os.path.join(args.output_dir, "pytorch_model.bin")
torch.save(model_to_save.state_dict(), output_model_file)
# Save a trained model
model_to_save = model.module if hasattr(model, 'module') else model # Only save the model it-self
torch.save(model_to_save.state_dict(), output_model_file)
# Load a trained model that you have fine-tuned
model_state_dict = torch.load(output_model_file)
model = BertForQuestionAnswering.from_pretrained(args.bert_model, state_dict=model_state_dict)
model.to(device)
if (args.do_predict or args.do_benchmark) and (args.local_rank == -1 or torch.distributed.get_rank() == 0):
# Load a trained model that you have fine-tuned
model_state_dict = torch.load(output_model_file) if os.path.exists(output_model_file) else None
model = BertForQuestionAnswering.from_pretrained(args.bert_model, state_dict=model_state_dict)
model.to(device)
if args.predict_fp16:
model.half()
if args.do_predict and (args.local_rank == -1 or torch.distributed.get_rank() == 0):
eval_examples = read_squad_examples(
input_file=args.predict_file, is_training=False)
eval_features = convert_examples_to_features(
@@ -953,28 +965,37 @@ def main():
model.eval()
all_results = []
logger.info("Start evaluating")
start_time = datetime.now()
for input_ids, input_mask, segment_ids, example_indices in tqdm(eval_dataloader, desc="Evaluating"):
if len(all_results) % 1000 == 0:
logger.info("Processing example: %d" % (len(all_results)))
input_ids = input_ids.to(device)
input_mask = input_mask.to(device)
segment_ids = segment_ids.to(device)
with torch.no_grad():
batch_start_logits, batch_end_logits = model(input_ids, segment_ids, input_mask)
for i, example_index in enumerate(example_indices):
start_logits = batch_start_logits[i].detach().cpu().tolist()
end_logits = batch_end_logits[i].detach().cpu().tolist()
eval_feature = eval_features[example_index.item()]
unique_id = int(eval_feature.unique_id)
all_results.append(RawResult(unique_id=unique_id,
start_logits=start_logits,
end_logits=end_logits))
output_prediction_file = os.path.join(args.output_dir, "predictions.json")
output_nbest_file = os.path.join(args.output_dir, "nbest_predictions.json")
write_predictions(eval_examples, eval_features, all_results,
args.n_best_size, args.max_answer_length,
args.do_lower_case, output_prediction_file,
output_nbest_file, args.verbose_logging)
if args.do_benchmark:
with torch.no_grad():
model(input_ids, segment_ids, input_mask)
else:
with torch.no_grad():
batch_start_logits, batch_end_logits = model(input_ids, segment_ids, input_mask)
for i, example_index in enumerate(example_indices):
start_logits = batch_start_logits[i].detach().cpu().tolist()
end_logits = batch_end_logits[i].detach().cpu().tolist()
eval_feature = eval_features[example_index.item()]
unique_id = int(eval_feature.unique_id)
all_results.append(RawResult(unique_id=unique_id,
start_logits=start_logits,
end_logits=end_logits))
total_time = datetime.now() - start_time
logger.info("Total time: %d seconds — Number of examples: %d — Examples/second: %f" % (
total_time.total_seconds(), len(eval_data), len(eval_data)/total_time.total_seconds()))
if not args.do_benchmark:
output_prediction_file = os.path.join(args.output_dir, "predictions.json")
output_nbest_file = os.path.join(args.output_dir, "nbest_predictions.json")
write_predictions(eval_examples, eval_features, all_results,
args.n_best_size, args.max_answer_length,
args.do_lower_case, output_prediction_file,
output_nbest_file, args.verbose_logging)
if __name__ == "__main__":