predict3

复制代码
# predict.py

import numpy as np
import torch
import torch.nn as nn
from sklearn.preprocessing import OneHotEncoder
from model import FCNet
import pickle
import os
import sklearn
from packaging import version


# 自定义函数用于加载模型和编码器
def load_model(model_path):
    if not os.path.exists(model_path):
        raise FileNotFoundError(f"模型文件 '{model_path}' 不存在。")

    # 加载模型文件
    checkpoint = torch.load(model_path, map_location=torch.device('cpu'))
    encoder = checkpoint['encoder']
    amino_acids = checkpoint['amino_acids']

    # 定义模型
    feature_dim1 = 3  # 最多三个氨基酸
    feature_dim2 = len(amino_acids) + 1  # one-hot + 数值特征
    model = FCNet(feature_dim1, feature_dim2)
    model.load_state_dict(checkpoint['model_state_dict'])
    model.eval()

    return model, encoder, amino_acids


# 自定义函数进行预测
def predict(model, encoder, amino_acids, single_dict, aa_input):
    # 检查输入类型和长度
    if not isinstance(aa_input, str):
        raise ValueError("输入的氨基酸组合必须是字符串,例如 'A', 'AC', 'ACD'。")

    aa_input = aa_input.upper()
    length = len(aa_input)
    if length == 0 or length > 3:
        raise ValueError("输入的氨基酸组合长度必须在1到3之间。")

    # 填充至三个氨基酸
    if length == 1:
        aa_triplet = aa_input + '--'
    elif length == 2:
        aa_triplet = aa_input + '-'
    else:
        aa_triplet = aa_input

    # 检查是否所有氨基酸都在编码器中
    for aa in aa_triplet:
        if aa not in amino_acids:
            raise ValueError(f"氨基酸 '{aa}' 不在编码器的氨基酸列表中。")

    # 独热编码
    try:
        one_hot1 = encoder.transform([[aa_triplet[0]]])[0]
        one_hot2 = encoder.transform([[aa_triplet[1]]])[0]
        one_hot3 = encoder.transform([[aa_triplet[2]]])[0]
    except Exception as e:
        raise ValueError(f"无法对氨基酸组合 '{aa_triplet}' 进行编码。错误信息:{e}")

    # 读取单个氨基酸的数值特征
    value1_aa1 = single_dict.get(aa_triplet[0], 0.0)
    value1_aa2 = single_dict.get(aa_triplet[1], 0.0)
    value1_aa3 = single_dict.get(aa_triplet[2], 0.0)

    # 构建特征向量
    feature_part1 = np.concatenate([one_hot1, [value1_aa1]])
    feature_part2 = np.concatenate([one_hot2, [value1_aa2]])
    feature_part3 = np.concatenate([one_hot3, [value1_aa3]])
    feature = np.stack([feature_part1, feature_part2, feature_part3])  # (3, N +1)

    # 转换为张量并添加 batch 维度
    feature_tensor = torch.tensor(feature, dtype=torch.float32).unsqueeze(0)  # (1, 3, N +1)

    # 进行预测
    with torch.no_grad():
        output = model(feature_tensor)

    # 获取预测结果
    predicted_values = output.squeeze(0).numpy()  # (3,)
    return predicted_values


def main():
    import pandas as pd  # 确保 pandas 已导入

    # 模型文件路径
    model_path = '../models/model.pth'

    # 加载模型、编码器和氨基酸列表
    try:
        model, encoder, amino_acids = load_model(model_path)
    except Exception as e:
        print(f"加载模型失败:{e}")
        return

    # 加载单个氨基酸的数值字典
    single_dict_path = '../data/single_dict.pkl'
    if not os.path.exists(single_dict_path):
        # 如果未保存,读取 single.csv 并创建字典,然后保存
        single_csv = '../data/single.csv'
        if not os.path.exists(single_csv):
            raise FileNotFoundError(f"单个氨基酸数据文件 '{single_csv}' 不存在。")
        single_df = pd.read_csv(single_csv, header=None, names=['AminoAcid', 'Value'])
        single_dict = single_df.set_index('AminoAcid')['Value'].to_dict()
        # 保存字典
        with open(single_dict_path, 'wb') as f:
            pickle.dump(single_dict, f)
    else:
        with open(single_dict_path, 'rb') as f:
            single_dict = pickle.load(f)

    # 提示用户输入氨基酸组合
    while True:
        aa_input = input("请输入氨基酸组合(例如 'A'、'AC'、'ACD'),输入 'exit' 退出:").strip().upper()
        if aa_input.lower() == 'exit':
            print("退出预测程序。")
            break
        try:
            predicted = predict(model, encoder, amino_acids, single_dict, aa_input)
            # 根据输入长度,打印对应数量的值
            length = len(aa_input)
            if length == 1:
                print(f"预测结果 - Value1: {predicted[0]:.6f}")
            elif length == 2:
                print(f"预测结果 - Value1: {predicted[0]:.6f}, Value2: {predicted[1]:.6f}")
            else:
                print(f"预测结果 - Value1: {predicted[0]:.6f}, Value2: {predicted[1]:.6f}, Value3: {predicted[2]:.6f}")
        except Exception as e:
            print(f"错误:{e}")


if __name__ == '__main__':
    main()
相关推荐
久久学姐1 小时前
Python+Playwright+Pytest+BDD,用FSM打造高效测试框架
python·pytest·bdd·playwright·fsm
Zane19942 小时前
闭包到底"闭"住了什么?一文讲透 LEGB 规则与循环里的闭包陷阱
后端·python
Lumi_Peak2 小时前
Claude思考了42秒,我的代码质量直接提升了一个档次
python·claude
m沐沐3 小时前
【机器学习】DBSCAN聚类算法——原理、参数调优与实战
人工智能·python·深度学习·算法·机器学习·聚类·dbscan
梅孔立3 小时前
推荐一个 Python 开源项目:AI 模板填充 + Markdown 转 Word,面向 Aspose 模板引擎的效率神器
人工智能·python·开源
码农小韩3 小时前
AIAgent应用开发——大模型理论基础与应用(七)
python·学习·ai·大模型·agent
CodeLinghu3 小时前
LangSmith Evaluate实战评估Agent
人工智能·python·语言模型·llm
起司喵喵4 小时前
推荐一款基于 Python 和 Rust 开发的跨平台 GUI 自动化库!
python·rust·自动化
过期的秋刀鱼!4 小时前
学习曲线-过拟合和欠拟合要做什么以及原因
人工智能·python·深度学习·算法·机器学习·模型评估
Yolanda_20225 小时前
Python学习-第九部分-错误处理与异常处理
开发语言·python·学习