TensorFlow 모델 직렬화 포맷 완벽 가이드

TensorFlow는 학습된 모델을 저장하고 배포하기 위해 여러 가지 직렬화 포맷을 제공합니다. 각 포맷은 특정 사용 사례에 최적화되어 있으며, 상호 변환도 가능합니다. 주요 포맷으로는 Checkpoint, GraphDef, SavedModel이 있습니다.

1. Checkpoint (*.ckpt)

학습 과정 중 중간 상태를 저장하는 데 사용되는 포맷입니다. tf.train.Saver 클래스를 통해 생성되며, 모델의 가중치(Variable) 값만을 바이너리 형태로 저장합니다.

중요한 특징은 그래프 구조 자체는 저장하지 않는다는 점입니다. 따라서 체크포인트를 복원하려면 원본 파이썬 코드로 네트워크 구조를 재정의해야 합니다.

import tensorflow as tf

# 모델 구조 정의 (필수)
input_layer = tf.placeholder(tf.float32, shape=[None, 784], name='input')
weights = tf.Variable(tf.random_normal([784, 10]), name='weights')
bias = tf.Variable(tf.zeros([10]), name='bias')
predictions = tf.add(tf.matmul(input_layer, weights), bias, name='output')

# 세션 초기화 및 저장
with tf.Session() as session:
    session.run(tf.global_variables_initializer())
    
    # 학습 수행...
    
    saver_handler = tf.train.Saver()
    saver_handler.save(session, './model_checkpoint/model.ckpt')

복원 시에도 동일한 구조 정의가 필요합니다:

# 동일한 구조 재정의 후
restorer = tf.train.Saver()
with tf.Session() as session:
    restorer.restore(session, './model_checkpoint/model.ckpt')
    # 이후 추론 수행

2. GraphDef (*.pb)

Protocol Buffer 형식으로 직렬화된 계산 그래프를 저장합니다. 연산 노드, 텐서 형태, 변수 정의 등 그래프의 완전한 구조를 포함하지만, 실제 학습된 가중치 값은 포함하지 않습니다.

GraphDef에서 그래프를 복원하는 방법:

from tensorflow.core.framework import graph_pb2

graph_definition = graph_pb2.GraphDef()

with open('network_structure.pb', 'rb') as pb_file:
    graph_definition.ParseFromString(pb_file.read())

tf.import_graph_def(graph_definition, name='')

Frozen GraphDef

배포 환경에서 자주 사용되는 변형 포맷입니다. 모든 Variable 노드를 상수(Constant)로 변환하여 가중치 값을 그래프 자체에 내장시킵니다. tensorflow/python/tools/freeze_graph.py 유틸리티로 생성할 수 있습니다.

Protobuf는 텍스트 형식(.pbtxt)도 지원하지만, 가중치가 포함된 경우 파일 크기가 급증하므로 바이너리 형식을 권장합니다.

3. SavedModel (권장 포맷)

Google이 공식적으로 권장하는 크로스 플랫폼 포맷입니다. 그래프 구조와 가중치를 통합 저장하고, 서명(Signature)을 통해 입출력 인터페이스를 명시적으로 정의합니다. 언어 독립적이며 TensorFlow Serving, TensorFlow.js, TensorFlow Lite 등 다양한 런타임에서 활용 가능합니다.

디렉터리 구조

saved_model_directory/
├── saved_model.pb          # MetaGraphDef (그래프 및 메타데이터)
├── variables/
│   ├── variables.data-00000-of-00001
│   └── variables.index     # 체크포인트 형식의 가중치
└── assets/                 # 추가 리소스 (어휘 사전 등)

저장 구현 예시

방식 A: 명시적 텐서 정보 구성

import tensorflow as tf

class ModelExporter:
    def __init__(self, feature_input, label_input, keep_prob, prediction_output):
        self.feature_tensor = feature_input
        self.label_tensor = label_input
        self.dropout_rate = keep_prob
        self.logits = prediction_output
    
    def create_serving_signature(self):
        input_signatures = {
            'features': tf.saved_model.utils.build_tensor_info(self.feature_tensor),
            'labels': tf.saved_model.utils.build_tensor_info(self.label_tensor),
            'dropout_keep': tf.saved_model.utils.build_tensor_info(self.dropout_rate)
        }
        
        output_signatures = {
            'class_predictions': tf.saved_model.utils.build_tensor_info(self.logits)
        }
        
        return tf.saved_model.signature_def_utils.build_signature_def(
            inputs=input_signatures,
            outputs=output_signatures,
            method_name=tf.saved_model.signature_constants.PREDICT_METHOD_NAME
        )

def persist_model(session, exporter, export_path):
    builder = tf.saved_model.builder.SavedModelBuilder(export_path)
    
    signature = exporter.create_serving_signature()
    
    builder.add_meta_graph_and_variables(
        session,
        tags=[tf.saved_model.tag_constants.SERVING],
        signature_def_map={'classify': signature},
        clear_devices=True
    )
    builder.save()

방식 B: 간결한 서명 정의

def generate_signature_wrapper(model_instance):
    input_map = {
        'features': model_instance.feature_tensor,
        'labels': model_instance.label_tensor,
        'dropout_keep': model_instance.dropout_rate
    }
    output_map = {'class_predictions': model_instance.logits}
    
    prediction_signature = tf.saved_model.signature_def_utils.predict_signature_def(
        inputs=input_map,
        outputs=output_map
    )
    
    return {
        tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY: 
        prediction_signature
    }

def export_with_wrapper(session, model, target_dir):
    builder = tf.saved_model.builder.SavedModelBuilder(target_dir)
    signatures = generate_signature_wrapper(model)
    
    builder.add_meta_graph_and_variables(
        sess=session,
        tags=[tf.saved_model.tag_constants.SERVING],
        signature_def_map=signatures,
        clear_devices=True
    )
    builder.save()

모델 로드 및 추론

import tensorflow as tf

def load_and_infer(export_directory, sample_features, sample_labels):
    with tf.Session(graph=tf.Graph()) as sess:
        # SavedModel 로드
        loaded_meta = tf.saved_model.loader.load(
            sess,
            [tf.saved_model.tag_constants.SERVING],
            export_directory
        )
        
        # 서명에서 텐서 이름 추출
        signature_key = 'classify'
        signature = loaded_meta.signature_def[signature_key]
        
        feature_tensor_name = signature.inputs['features'].name
        label_tensor_name = signature.inputs['labels'].name
        dropout_tensor_name = signature.inputs['dropout_keep'].name
        output_tensor_name = signature.outputs['class_predictions'].name
        
        # 실제 텐서 객체 획득
        feature_placeholder = sess.graph.get_tensor_by_name(feature_tensor_name)
        label_placeholder = sess.graph.get_tensor_by_name(label_tensor_name)
        dropout_placeholder = sess.graph.get_tensor_by_name(dropout_tensor_name)
        output_tensor = sess.graph.get_tensor_by_name(output_tensor_name)
        
        # 추론 실행
        results = sess.run(
            output_tensor,
            feed_dict={
                feature_placeholder: sample_features,
                label_placeholder: sample_labels,
                dropout_placeholder: 1.0  # 추론 시 드롭아웃 비활성화
            }
        )
        return results

4. 포맷 선택 가이드

사용 사례권장 포맷이유
학습 중단점 저장Checkpoint빠른 저장/복원, 구조 유연성
그래프 구조 분석GraphDef시각화 및 변환 도구 호환
프로덕션 서빙SavedModel버전 관리, 서명 기반 추론
모바일/임베디드 배포Frozen GraphDef / TFLite최소 의존성, 경량화

5. 상호 변환 개요

TensorFlow는 포맷 간 변환을 위한 공식 도구를 제공합니다:

  • Checkpoint → SavedModel: tf.saved_model.builder API 활용
  • SavedModel → Frozen GraphDef: freeze_graph 도구 사용
  • Any → TensorFlow Lite: TFLiteConverter를 통한 최적화 변환

각 변환 과정에서 양자화(quantization)나 프루닝(pruning) 같은 최적화 기법을 적용할 수 있으며, 이는 특히 엣지 디바이스 배포 시 필수적입니다.

태그: TensorFlow SavedModel GraphDef Checkpoint Model Serialization

10월 1일 05:27에 게시됨