Python AI基础:文件操作与数据持久化
引言
在人工智能和机器学习项目中,数据是核心驱动力。无论是训练模型、加载预训练权重,还是保存中间结果,都离不开高效、可靠的文件操作与数据持久化技术。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. 最佳实践与总结
-
路径处理:使用
os.path.join()或pathlib.Path来构建跨平台兼容的文件路径。from pathlib import Path data_dir = Path('./data') model_path = data_dir / 'models' / 'final_model.h5' -
上下文管理器:始终使用
with open(...) as f:语句来确保文件被正确关闭,即使在发生异常时也是如此。 -
版本控制与命名:为重要的模型和数据集文件使用有意义的、包含版本信息的命名(如
model_v2.1_20250705.pt),并考虑使用工具(如 DVC)进行数据版本管理。 -
安全性:谨慎处理来自外部的文件,特别是反序列化操作(如
pickle.load,torch.load)。优先使用更安全的格式如 JSON。 -
效率:对于大型数值数据集,优先使用
.npy,.npz,.parquet或 HDF5 等二进制格式,它们比文本格式读写更快、体积更小。
总结:文件操作与数据持久化是连接 AI 算法与现实世界的桥梁。从简单的文本配置到复杂的模型权重,从临时的实验记录到可追溯的数据库,选择正确的工具和方法,能让你的 AI 项目更加健壮、高效和可维护。建议在实际项目中根据数据规模、读写频率和协作需求,灵活组合运用上述技术。
更多推荐




所有评论(0)