引言

在人工智能和机器学习项目中,数据是核心驱动力。无论是训练模型、加载预训练权重,还是保存中间结果,都离不开高效、可靠的文件操作与数据持久化技术。Python 作为 AI 领域的主流语言,提供了丰富而强大的内置模块和第三方库来处理这些任务。掌握这些技能,是构建稳定、可复现 AI 应用的基础。本文将系统介绍 Python 中用于文件操作与数据持久化的核心方法,并结合 AI 场景下的典型用例进行讲解。

1. 基础文件操作

Python 内置的 open() 函数是文件操作的起点。它支持多种模式,以适应不同的读写需求。

1.1 读写文本文件

文本文件是存储配置、日志和简单数据集的常见格式。

# 写入文本文件
with open('config.txt', 'w', encoding='utf-8') as f:
    f.write('模型名称: resnet50\n')
    f.write('学习率: 0.001\n')
    f.write('训练轮数: 100\n')

# 读取文本文件
with open('config.txt', 'r', encoding='utf-8') as f:
    content = f.read()
    print(content)
# 输出:
# 模型名称: resnet50
# 学习率: 0.001
# 训练轮数: 100

# 逐行读取(适用于大文件)
with open('dataset_labels.txt', 'r') as f:
    for line in f:
        label = line.strip()  # 去除换行符
        # 处理每一行标签

AI 应用场景:读取训练数据集的标签文件、保存模型训练的超参数配置、记录训练过程的日志信息。

1.2 读写二进制文件

二进制模式用于处理非文本数据,如图片、音频、视频以及序列化的模型权重。

# 写入二进制数据(例如,保存一个 NumPy 数组)
import numpy as np
data_array = np.random.randn(100, 10).astype(np.float32)
with open('sample_data.bin', 'wb') as f:
    f.write(data_array.tobytes())  # 将数组转换为字节写入

# 读取二进制数据
with open('sample_data.bin', 'rb') as f:
    bytes_data = f.read()
    # 将字节重新转换为 NumPy 数组
    loaded_array = np.frombuffer(bytes_data, dtype=np.float32).reshape(100, 10)

AI 应用场景:加载原始的图像或音频数据、保存和加载自定义的二进制格式数据集。

2. 结构化数据持久化

对于更复杂的数据结构(如列表、字典、自定义对象),我们需要序列化和反序列化。

2.1 使用 pickle 模块

pickle 是 Python 标准的序列化模块,可以将几乎任何 Python 对象转换为字节流。

import pickle

# 要保存的数据(可以是复杂的嵌套结构)
model_metadata = {
    'name': 'MyClassifier',
    'version': '1.0',
    'hyperparameters': {'lr': 0.01, 'batch_size': 32},
    'class_labels': ['cat', 'dog', 'bird']
}

# 序列化并保存到文件
with open('model_meta.pkl', 'wb') as f:
    pickle.dump(model_metadata, f)

# 从文件加载并反序列化
with open('model_meta.pkl', 'rb') as f:
    loaded_metadata = pickle.load(f)
    print(loaded_metadata['hyperparameters'])  # 输出: {'lr': 0.01, 'batch_size': 32}

警告pickle 文件可能包含恶意代码,只应加载来自可信来源的文件。

AI 应用场景:保存 Scikit-learn 训练好的模型、存储复杂的数据预处理管道、缓存中间计算结果。

2.2 使用 json 模块

JSON (JavaScript Object Notation) 是一种轻量级、跨语言的数据交换格式。它比 pickle 更安全,但只能处理基本数据类型(字典、列表、字符串、数字、布尔值和 None)。

import json

# 将 Python 对象转换为 JSON 字符串并保存
training_log = {
    'epoch': [1, 2, 3, 4, 5],
    'train_loss': [0.5, 0.3, 0.2, 0.15, 0.12],
    'val_accuracy': [0.85, 0.88, 0.90, 0.91, 0.92]
}

with open('training_log.json', 'w', encoding='utf-8') as f:
    json.dump(training_log, f, indent=4)  # indent 参数使文件更易读

# 从 JSON 文件读取并解析为 Python 对象
with open('training_log.json', 'r', encoding='utf-8') as f:
    loaded_log = json.load(f)
    print(f"最终验证准确率: {loaded_log['val_accuracy'][-1]}")  # 输出: 最终验证准确率: 0.92

AI 应用场景:保存模型配置、记录实验指标、与前端或其他服务交换预测结果。

3. 科学计算与 AI 专用格式

在 AI 领域,一些专门为高效存储和读取数值数据而设计的库和格式被广泛使用。

3.1 NumPy (*.npy, *.npz)

NumPy 提供了专为数组数据优化的存储格式。

import numpy as np

# 保存单个数组到 .npy 文件
weights = np.random.randn(256, 128)
np.save('model_weights.npy', weights)

# 加载 .npy 文件
loaded_weights = np.load('model_weights.npy')
print(loaded_weights.shape)  # 输出: (256, 128)

# 保存多个数组到 .npz 文件(压缩格式)
features = np.random.randn(1000, 50)
labels = np.random.randint(0, 10, size=1000)
np.savez('dataset.npz', X=features, y=labels)

# 加载 .npz 文件
data = np.load('dataset.npz')
print(data['X'].shape, data['y'].shape)  # 输出: (1000, 50) (1000,)

3.2 PyTorch (*.pt, *.pth)

PyTorch 使用 torch.save()torch.load() 来保存和加载模型的状态字典或整个模型。

import torch
import torch.nn as nn

# 定义一个简单的模型
class SimpleNN(nn.Module):
    def __init__(self):
        super().__init__()
        self.fc = nn.Linear(10, 2)
    def forward(self, x):
        return self.fc(x)

model = SimpleNN()

# 保存模型的“状态字典”(推荐方式,只保存可学习参数)
torch.save(model.state_dict(), 'model_state_dict.pt')

# 加载状态字典到模型(需要先实例化一个结构相同的模型)
new_model = SimpleNN()
new_model.load_state_dict(torch.load('model_state_dict.pt'))

# 也可以保存整个模型(包含结构,但可能更不灵活)
torch.save(model, 'entire_model.pth')
loaded_entire_model = torch.load('entire_model.pth')

3.3 TensorFlow / Keras (*.h5, SavedModel)

Keras 提供了简单易用的模型保存与加载 API。

import tensorflow as tf

# 假设 `model` 是一个已编译的 Keras 模型
# model = tf.keras.Sequential([...])

# 保存为 HDF5 格式(.h5)
model.save('my_keras_model.h5')

# 加载模型
loaded_model = tf.keras.models.load_model('my_keras_model.h5')

# 保存为 TensorFlow SavedModel 格式(目录)
model.save('my_saved_model/')  # 会创建一个目录

# 加载 SavedModel
loaded_saved_model = tf.keras.models.load_model('my_saved_model/')

4. 高级持久化与数据库

对于大规模、需要快速查询的数据,或需要持久化复杂关系的数据,数据库是更好的选择。

4.1 SQLite(轻量级数据库)

Python 内置 sqlite3 模块,无需安装额外服务。

import sqlite3
import pandas as pd

# 连接到数据库(如果不存在则创建)
conn = sqlite3.connect('experiments.db')
cursor = conn.cursor()

# 创建表
cursor.execute('''
CREATE TABLE IF NOT EXISTS training_runs (
    id INTEGER PRIMARY KEY AUTOINCREMENT,
    model_name TEXT NOT NULL,
    accuracy REAL,
    created_date TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
''')

# 插入实验记录
cursor.execute("INSERT INTO training_runs (model_name, accuracy) VALUES (?, ?)", ('ResNet50', 0.945))
conn.commit()

# 查询数据
df = pd.read_sql_query("SELECT * FROM training_runs", conn)
print(df)

conn.close()

4.2 使用 Pandas 直接读写

Pandas 可以方便地将 DataFrame 与多种文件格式互转。

import pandas as pd

# 创建示例数据
df = pd.DataFrame({
    'image_path': ['img1.jpg', 'img2.jpg', 'img3.jpg'],
    'label': [0, 1, 0],
    'confidence': [0.99, 0.87, 0.92]
})

# 保存为 CSV
df.to_csv('predictions.csv', index=False)

# 保存为 Parquet(列式存储,高效压缩)
df.to_parquet('predictions.parquet', index=False)

# 读取文件
df_from_csv = pd.read_csv('predictions.csv')
df_from_parquet = pd.read_parquet('predictions.parquet')

5. 最佳实践与总结

  1. 路径处理:使用 os.path.join()pathlib.Path 来构建跨平台兼容的文件路径。

    from pathlib import Path
    data_dir = Path('./data')
    model_path = data_dir / 'models' / 'final_model.h5'
    
  2. 上下文管理器:始终使用 with open(...) as f: 语句来确保文件被正确关闭,即使在发生异常时也是如此。

  3. 版本控制与命名:为重要的模型和数据集文件使用有意义的、包含版本信息的命名(如 model_v2.1_20250705.pt),并考虑使用工具(如 DVC)进行数据版本管理。

  4. 安全性:谨慎处理来自外部的文件,特别是反序列化操作(如 pickle.load, torch.load)。优先使用更安全的格式如 JSON。

  5. 效率:对于大型数值数据集,优先使用 .npy, .npz, .parquet 或 HDF5 等二进制格式,它们比文本格式读写更快、体积更小。

总结:文件操作与数据持久化是连接 AI 算法与现实世界的桥梁。从简单的文本配置到复杂的模型权重,从临时的实验记录到可追溯的数据库,选择正确的工具和方法,能让你的 AI 项目更加健壮、高效和可维护。建议在实际项目中根据数据规模、读写频率和协作需求,灵活组合运用上述技术。

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐