You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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.py and navigate to the end of the training flow (right after tf.train.train_and_evaluate or 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_config matches the configuration used during pre-training. After training completes, you'll find the full SavedModel structure in the saved_model subfolder 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:

  1. Gather your bert_config.json (same as used in pre-training) and the directory containing your model.ckpt files.
  2. 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:

  1. Install dependencies:
    pip install transformers datasets
    
  2. 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.
  3. Run the run_mlm.py script (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-uncased and --tokenizer_name bert-base-uncased to 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.29 02:52:38