TensorFlow官方实现预训练自定义数据后BERT模型保存及PyTorch兼容方案咨询
Hey there, let's break down your three questions one by one with practical, tested solutions:
1. Save directly to .pb (SavedModel) after pre-training with official BERT run_pretraining.py
The official run_pretraining.py is built on TensorFlow 1.x, so we can add code to export the model to SavedModel format (the structure with assets, variables, and saved_model.pb) right after training finishes. Here's how:
- Open
run_pretraining.pyand navigate to the end of the training flow (right aftertf.train.train_and_evaluateor the training loop concludes). - Insert this code block to handle the export:
import os import tensorflow as tf from modeling import BertModel # Define input placeholders matching BERT's expected inputs input_ids = tf.placeholder(tf.int32, shape=[None, None], name="input_ids") input_mask = tf.placeholder(tf.int32, shape=[None, None], name="input_mask") segment_ids = tf.placeholder(tf.int32, shape=[None, None], name="segment_ids") # Reconstruct BERT model in inference mode model = BertModel( config=bert_config, is_training=False, input_ids=input_ids, input_mask=input_mask, token_type_ids=segment_ids, use_one_hot_embeddings=False ) # Load the trained checkpoint saver = tf.train.Saver() with tf.Session() as sess: saver.restore(sess, tf.train.latest_checkpoint(FLAGS.output_dir)) # Define serving signature (adjust outputs based on your needs) signature_def = tf.saved_model.signature_def_utils.predict_signature_def( inputs={ "input_ids": input_ids, "input_mask": input_mask, "segment_ids": segment_ids }, outputs={ "pooled_output": model.get_pooled_output(), "sequence_output": model.get_sequence_output() } ) # Save the SavedModel save_path = os.path.join(FLAGS.output_dir, "saved_model") builder = tf.saved_model.builder.SavedModelBuilder(save_path) builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={ tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: signature_def } ) builder.save() - Make sure you have all required imports, and
bert_configmatches the configuration used during pre-training. After training completes, you'll find the full SavedModel structure in thesaved_modelsubfolder of your output directory.
2. Convert existing model.ckpt files to .pb (SavedModel)
If you already have trained checkpoints and don't want to re-run training, create a standalone conversion script:
- Gather your
bert_config.json(same as used in pre-training) and the directory containing yourmodel.ckptfiles. - Use this script to handle the conversion:
import tensorflow as tf from modeling import BertModel import os # Configure paths CKPT_DIR = "/path/to/your/ckpt-directory" OUTPUT_DIR = "/path/to/save/saved_model" BERT_CONFIG_PATH = "/path/to/bert_config.json" # Load BERT config bert_config = modeling.BertConfig.from_json_file(BERT_CONFIG_PATH) # Reset graph and build inference model tf.reset_default_graph() input_ids = tf.placeholder(tf.int32, shape=[None, None], name="input_ids") input_mask = tf.placeholder(tf.int32, shape=[None, None], name="input_mask") segment_ids = tf.placeholder(tf.int32, shape=[None, None], name="segment_ids") model = BertModel( config=bert_config, is_training=False, input_ids=input_ids, input_mask=input_mask, token_type_ids=segment_ids, use_one_hot_embeddings=False ) # Load checkpoint and save SavedModel saver = tf.train.Saver() with tf.Session() as sess: latest_ckpt = tf.train.latest_checkpoint(CKPT_DIR) saver.restore(sess, latest_ckpt) # Define serving signature signature_def = tf.saved_model.signature_def_utils.predict_signature_def( inputs={"input_ids": input_ids, "input_mask": input_mask, "segment_ids": segment_ids}, outputs={"pooled_output": model.get_pooled_output(), "sequence_output": model.get_sequence_output()} ) builder = tf.saved_model.builder.SavedModelBuilder(OUTPUT_DIR) builder.add_meta_graph_and_variables( sess, [tf.saved_model.tag_constants.SERVING], signature_def_map={ tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: signature_def } ) builder.save()
- Run the script, and you'll get the full SavedModel structure in your specified output directory. Note: Use the same TensorFlow 1.x version as your pre-training to avoid compatibility issues.
3. Pre-train BERT directly in PyTorch for downstream tasks
This is the most straightforward approach if you need PyTorch compatibility long-term. The Hugging Face Transformers library provides an easy-to-use pre-training script for masked language modeling (MLM), which is the standard pre-training task for BERT:
- Install dependencies:
pip install transformers datasets - Prepare your custom data: Format it as a text file with one sample per line, or use JSONL. No extra annotations are needed for MLM pre-training.
- Run the
run_mlm.pyscript (included with Transformers) with a command like this:python run_mlm.py \ --model_name_or_path bert-base-uncased \ --train_file /path/to/your/train.txt \ --validation_file /path/to/your/val.txt \ --per_device_train_batch_size 16 \ --per_device_eval_batch_size 16 \ --do_train \ --do_eval \ --output_dir /path/to/save/pytorch-bert-model \ --num_train_epochs 3 \ --logging_steps 100 \ --save_steps 1000
- If you want to pre-train from scratch (not fine-tune
bert-base-uncased), set--model_name_or_path None, then add--config_name bert-base-uncasedand--tokenizer_name bert-base-uncasedto initialize the model structure and tokenizer. - After training, the output directory will contain a PyTorch-native model that you can load directly with:
from transformers import BertModel model = BertModel.from_pretrained("/path/to/save/pytorch-bert-model") - This model works seamlessly with all PyTorch downstream tasks, no conversion needed.
内容的提问来源于stack exchange,提问作者popeye

