FScanpy 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- FScanpy/__init__.py +206 -0
- FScanpy/data/__init__.py +122 -0
- FScanpy/data/test_data/blastx_example.xlsx +0 -0
- FScanpy/data/test_data/full_seq.xlsx +0 -0
- FScanpy/data/test_data/mrna_example.fasta +2584 -0
- FScanpy/data/test_data/region_example.csv +4 -0
- FScanpy/features/__init__.py +4 -0
- FScanpy/features/cnn_input.py +79 -0
- FScanpy/features/sequence.py +283 -0
- FScanpy/predictor.py +616 -0
- FScanpy/pretrained/long.pth +4 -0
- FScanpy/pretrained/short.pkl +0 -0
- FScanpy/utils.py +203 -0
- fscanpy-1.0.0.dist-info/METADATA +262 -0
- fscanpy-1.0.0.dist-info/RECORD +18 -0
- fscanpy-1.0.0.dist-info/WHEEL +5 -0
- fscanpy-1.0.0.dist-info/licenses/LICENSE +21 -0
- fscanpy-1.0.0.dist-info/top_level.txt +1 -0
FScanpy/predictor.py
ADDED
|
@@ -0,0 +1,616 @@
|
|
|
1
|
+
import os
|
|
2
|
+
from pathlib import Path
|
|
3
|
+
import pickle
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
import torch
|
|
7
|
+
import torch.nn as nn
|
|
8
|
+
from .features.sequence import SequenceFeatureExtractor
|
|
9
|
+
from .features.cnn_input import CNNInputProcessor
|
|
10
|
+
from .utils import extract_window_sequences
|
|
11
|
+
import matplotlib.pyplot as plt
|
|
12
|
+
import joblib
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class PRFPredictor:
|
|
16
|
+
|
|
17
|
+
def __init__(self, model_dir=None):
|
|
18
|
+
"""
|
|
19
|
+
初始化PRF预测器
|
|
20
|
+
|
|
21
|
+
Args:
|
|
22
|
+
model_dir: 模型目录路径(可选)
|
|
23
|
+
"""
|
|
24
|
+
if model_dir is None:
|
|
25
|
+
model_dir = Path(__file__).resolve().parent / 'pretrained'
|
|
26
|
+
|
|
27
|
+
try:
|
|
28
|
+
# 设备
|
|
29
|
+
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
|
|
30
|
+
|
|
31
|
+
# 加载模型 - 使用新的命名约定
|
|
32
|
+
self.short_model = self._load_pickle(os.path.join(model_dir, 'short.pkl')) # HistGB模型
|
|
33
|
+
|
|
34
|
+
# 优先使用 PyTorch 权重 long.pth;若不存在则回退到 long.pkl(兼容旧版本)
|
|
35
|
+
long_pth = os.path.join(model_dir, 'long.pth')
|
|
36
|
+
long_pkl = os.path.join(model_dir, 'long.pkl')
|
|
37
|
+
if os.path.exists(long_pth):
|
|
38
|
+
self.long_model = self._load_long_torch(long_pth)
|
|
39
|
+
else:
|
|
40
|
+
self.long_model = self._load_pickle(long_pkl)
|
|
41
|
+
|
|
42
|
+
# 初始化特征提取器和CNN处理器,使用与训练时相同的序列长度
|
|
43
|
+
self.short_seq_length = 33 # HistGB使用的序列长度
|
|
44
|
+
self.long_seq_length = 399 # BiLSTM-CNN使用的序列长度
|
|
45
|
+
|
|
46
|
+
# 初始化特征提取器和CNN输入处理器
|
|
47
|
+
self.feature_extractor = SequenceFeatureExtractor(seq_length=self.short_seq_length)
|
|
48
|
+
self.cnn_processor = CNNInputProcessor(max_length=self.long_seq_length)
|
|
49
|
+
|
|
50
|
+
# 检测模型类型以优化预测性能
|
|
51
|
+
self._detect_model_types()
|
|
52
|
+
|
|
53
|
+
except FileNotFoundError as e:
|
|
54
|
+
raise FileNotFoundError(f"无法找到模型文件: {str(e)}。请确保 'short.pkl' 存在,且优先使用 'long.pth'(若无则需要 'long.pkl') 位于 {model_dir}")
|
|
55
|
+
except Exception as e:
|
|
56
|
+
raise Exception(f"加载模型出错: {str(e)}")
|
|
57
|
+
|
|
58
|
+
def _load_pickle(self, path):
|
|
59
|
+
"""安全加载pickle文件"""
|
|
60
|
+
try:
|
|
61
|
+
return joblib.load(path)
|
|
62
|
+
except Exception as e:
|
|
63
|
+
raise FileNotFoundError(f"无法加载模型文件 {path}: {str(e)}")
|
|
64
|
+
|
|
65
|
+
# ===== PyTorch Long 模型实现(与 internal_train_pytorch.py 对齐) =====
|
|
66
|
+
class _BiLSTM_CNN_Model(nn.Module):
|
|
67
|
+
def __init__(self, input_dim=1, embedding_dim=64, lstm_units=64, cnn_filters=64,
|
|
68
|
+
kernel_sizes=[3, 5, 7], dropout_rate=0.5, sequence_length=399):
|
|
69
|
+
super(PRFPredictor._BiLSTM_CNN_Model, self).__init__()
|
|
70
|
+
self.sequence_length = sequence_length
|
|
71
|
+
self.conv_layers = nn.ModuleList()
|
|
72
|
+
for kernel_size in kernel_sizes:
|
|
73
|
+
conv = nn.Sequential(
|
|
74
|
+
nn.Conv1d(in_channels=input_dim, out_channels=cnn_filters,
|
|
75
|
+
kernel_size=kernel_size, padding=kernel_size//2),
|
|
76
|
+
nn.BatchNorm1d(cnn_filters),
|
|
77
|
+
nn.ReLU(),
|
|
78
|
+
nn.MaxPool1d(kernel_size=2)
|
|
79
|
+
)
|
|
80
|
+
self.conv_layers.append(conv)
|
|
81
|
+
cnn_output_length = sequence_length // 2
|
|
82
|
+
total_cnn_features = len(kernel_sizes) * cnn_filters * cnn_output_length
|
|
83
|
+
self.lstm1 = nn.LSTM(input_dim, lstm_units, batch_first=True, bidirectional=True)
|
|
84
|
+
self.bn_lstm1 = nn.BatchNorm1d(sequence_length)
|
|
85
|
+
self.lstm2 = nn.LSTM(lstm_units * 2, lstm_units // 2, batch_first=True, bidirectional=True)
|
|
86
|
+
self.bn_lstm2 = nn.BatchNorm1d(lstm_units)
|
|
87
|
+
total_features = lstm_units + total_cnn_features
|
|
88
|
+
self.fc1 = nn.Linear(total_features, 256)
|
|
89
|
+
self.bn_fc1 = nn.BatchNorm1d(256)
|
|
90
|
+
self.dropout = nn.Dropout(dropout_rate)
|
|
91
|
+
self.fc2 = nn.Linear(256, 1)
|
|
92
|
+
def forward(self, x):
|
|
93
|
+
batch_size = x.size(0)
|
|
94
|
+
x_cnn = x.permute(0, 2, 1)
|
|
95
|
+
cnn_outputs = []
|
|
96
|
+
for conv in self.conv_layers:
|
|
97
|
+
conv_out = conv(x_cnn)
|
|
98
|
+
conv_out = conv_out.view(batch_size, -1)
|
|
99
|
+
cnn_outputs.append(conv_out)
|
|
100
|
+
cnn_merged = torch.cat(cnn_outputs, dim=1)
|
|
101
|
+
lstm_out, _ = self.lstm1(x)
|
|
102
|
+
lstm_out = self.bn_lstm1(lstm_out)
|
|
103
|
+
lstm_out, _ = self.lstm2(lstm_out)
|
|
104
|
+
lstm_out = lstm_out[:, -1, :]
|
|
105
|
+
lstm_out = self.bn_lstm2(lstm_out)
|
|
106
|
+
merged = torch.cat([lstm_out, cnn_merged], dim=1)
|
|
107
|
+
out = self.fc1(merged)
|
|
108
|
+
out = self.bn_fc1(out)
|
|
109
|
+
out = torch.relu(out)
|
|
110
|
+
out = self.dropout(out)
|
|
111
|
+
out = self.fc2(out)
|
|
112
|
+
out = torch.sigmoid(out)
|
|
113
|
+
return out.squeeze()
|
|
114
|
+
|
|
115
|
+
def _load_long_torch(self, checkpoint_path):
|
|
116
|
+
"""加载 PyTorch long 模型权重"""
|
|
117
|
+
model = PRFPredictor._BiLSTM_CNN_Model(
|
|
118
|
+
input_dim=1,
|
|
119
|
+
embedding_dim=64,
|
|
120
|
+
lstm_units=64,
|
|
121
|
+
cnn_filters=64,
|
|
122
|
+
kernel_sizes=[3, 5, 7],
|
|
123
|
+
dropout_rate=0.5,
|
|
124
|
+
sequence_length=399
|
|
125
|
+
).to(self.device)
|
|
126
|
+
# 兼容在不同设备保存/加载
|
|
127
|
+
state = torch.load(checkpoint_path, map_location=self.device, weights_only=True)
|
|
128
|
+
# 兼容 weights_only=True/False 的保存
|
|
129
|
+
if isinstance(state, dict) and all(k.startswith('module.') for k in state.keys()):
|
|
130
|
+
# 去掉分布式前缀
|
|
131
|
+
state = {k.replace('module.', '', 1): v for k, v in state.items()}
|
|
132
|
+
model.load_state_dict(state)
|
|
133
|
+
model.eval()
|
|
134
|
+
return model
|
|
135
|
+
|
|
136
|
+
def _detect_model_types(self):
|
|
137
|
+
"""检测模型类型以优化预测性能"""
|
|
138
|
+
self.short_is_sklearn = hasattr(self.short_model, 'predict_proba')
|
|
139
|
+
self.long_is_sklearn = hasattr(self.long_model, 'predict_proba')
|
|
140
|
+
try:
|
|
141
|
+
import torch as _t
|
|
142
|
+
self.long_is_torch = isinstance(self.long_model, nn.Module)
|
|
143
|
+
except Exception:
|
|
144
|
+
self.long_is_torch = False
|
|
145
|
+
|
|
146
|
+
def _predict_model(self, model, features, is_sklearn, seq_length):
|
|
147
|
+
"""统一的模型预测方法"""
|
|
148
|
+
try:
|
|
149
|
+
if is_sklearn:
|
|
150
|
+
# sklearn模型使用特征向量
|
|
151
|
+
if isinstance(features, np.ndarray) and features.ndim > 1:
|
|
152
|
+
features = features.flatten()
|
|
153
|
+
features_2d = np.array([features])
|
|
154
|
+
pred = model.predict_proba(features_2d)
|
|
155
|
+
return pred[0][1]
|
|
156
|
+
else:
|
|
157
|
+
# 深度学习模型(Keras 旧分支)
|
|
158
|
+
# 保留向后兼容,但 long 模型若为 torch 将不走此分支
|
|
159
|
+
if seq_length == self.long_seq_length:
|
|
160
|
+
model_input = self.cnn_processor.prepare_sequence(features)
|
|
161
|
+
else:
|
|
162
|
+
base_to_num = {'A': 1, 'T': 2, 'G': 3, 'C': 4, 'N': 0}
|
|
163
|
+
seq_numeric = [base_to_num.get(base, 0) for base in features.upper()]
|
|
164
|
+
model_input = np.array(seq_numeric).reshape(1, len(seq_numeric), 1)
|
|
165
|
+
try:
|
|
166
|
+
pred = model.predict(model_input, verbose=0)
|
|
167
|
+
except TypeError:
|
|
168
|
+
pred = model.predict(model_input)
|
|
169
|
+
if isinstance(pred, list):
|
|
170
|
+
pred = pred[0]
|
|
171
|
+
if hasattr(pred, 'shape') and len(pred.shape) > 1 and pred.shape[1] > 1:
|
|
172
|
+
return pred[0][1]
|
|
173
|
+
else:
|
|
174
|
+
return pred[0][0] if hasattr(pred[0], '__getitem__') else pred[0]
|
|
175
|
+
|
|
176
|
+
except Exception as e:
|
|
177
|
+
raise Exception(f"模型预测失败: {str(e)}")
|
|
178
|
+
|
|
179
|
+
def predict_single_position(self, fs_period, full_seq, short_threshold=0.1, ensemble_weight=0.4):
|
|
180
|
+
'''
|
|
181
|
+
预测单个位置的PRF状态
|
|
182
|
+
|
|
183
|
+
Args:
|
|
184
|
+
fs_period: 33bp序列 (short模型使用)
|
|
185
|
+
full_seq: 完整序列 (long模型使用)
|
|
186
|
+
short_threshold: short模型的概率阈值 (默认为0.1)
|
|
187
|
+
ensemble_weight: short模型在集成中的权重 (默认为0.4,long权重为0.6)
|
|
188
|
+
Returns:
|
|
189
|
+
dict: 包含预测概率的字典
|
|
190
|
+
'''
|
|
191
|
+
try:
|
|
192
|
+
# 验证权重参数
|
|
193
|
+
if not (0.0 <= ensemble_weight <= 1.0):
|
|
194
|
+
raise ValueError("ensemble_weight 必须在 0.0 到 1.0 之间")
|
|
195
|
+
|
|
196
|
+
long_weight = 1.0 - ensemble_weight
|
|
197
|
+
|
|
198
|
+
# 处理序列长度
|
|
199
|
+
if len(fs_period) > self.short_seq_length:
|
|
200
|
+
fs_period = self.feature_extractor.trim_sequence(fs_period, self.short_seq_length)
|
|
201
|
+
|
|
202
|
+
# Short模型预测 (HistGB)
|
|
203
|
+
try:
|
|
204
|
+
if self.short_is_sklearn:
|
|
205
|
+
short_features = self.feature_extractor.extract_features(fs_period)
|
|
206
|
+
short_prob = self._predict_model(self.short_model, short_features, True, self.short_seq_length)
|
|
207
|
+
else:
|
|
208
|
+
short_prob = self._predict_model(self.short_model, fs_period, False, self.short_seq_length)
|
|
209
|
+
except Exception as e:
|
|
210
|
+
print(f"Short模型预测时出错: {str(e)}")
|
|
211
|
+
short_prob = 0.0
|
|
212
|
+
|
|
213
|
+
# 如果short概率低于阈值,则跳过long模型
|
|
214
|
+
if short_prob < short_threshold:
|
|
215
|
+
return {
|
|
216
|
+
'Short_Probability': short_prob,
|
|
217
|
+
'Long_Probability': 0.0,
|
|
218
|
+
'Ensemble_Probability': 0.0,
|
|
219
|
+
'Ensemble_Weights': f'Short:{ensemble_weight:.1f}, Long:{long_weight:.1f}'
|
|
220
|
+
}
|
|
221
|
+
|
|
222
|
+
# Long模型预测 (BiLSTM-CNN)
|
|
223
|
+
try:
|
|
224
|
+
if getattr(self, 'long_is_torch', False):
|
|
225
|
+
long_prob = self._predict_long_torch(full_seq)
|
|
226
|
+
elif self.long_is_sklearn:
|
|
227
|
+
long_features = self.feature_extractor.extract_features(full_seq)
|
|
228
|
+
long_prob = self._predict_model(self.long_model, long_features, True, self.long_seq_length)
|
|
229
|
+
else:
|
|
230
|
+
long_prob = self._predict_model(self.long_model, full_seq, False, self.long_seq_length)
|
|
231
|
+
except Exception as e:
|
|
232
|
+
print(f"Long模型预测时出错: {str(e)}")
|
|
233
|
+
long_prob = 0.0
|
|
234
|
+
|
|
235
|
+
# 计算集成概率
|
|
236
|
+
try:
|
|
237
|
+
ensemble_prob = ensemble_weight * short_prob + long_weight * long_prob
|
|
238
|
+
except Exception as e:
|
|
239
|
+
print(f"计算集成概率时出错: {str(e)}")
|
|
240
|
+
ensemble_prob = (short_prob + long_prob) / 2
|
|
241
|
+
|
|
242
|
+
return {
|
|
243
|
+
'Short_Probability': short_prob,
|
|
244
|
+
'Long_Probability': long_prob,
|
|
245
|
+
'Ensemble_Probability': ensemble_prob,
|
|
246
|
+
'Ensemble_Weights': f'Short:{ensemble_weight:.1f}, Long:{long_weight:.1f}'
|
|
247
|
+
}
|
|
248
|
+
|
|
249
|
+
except Exception as e:
|
|
250
|
+
raise Exception(f"预测过程出错: {str(e)}")
|
|
251
|
+
|
|
252
|
+
# ===== Torch long 预测路径 =====
|
|
253
|
+
@staticmethod
|
|
254
|
+
def _process_sequence(seq):
|
|
255
|
+
seq = str(seq).upper()
|
|
256
|
+
return ''.join('N' if base not in 'ATCG' else base for base in seq)
|
|
257
|
+
|
|
258
|
+
@staticmethod
|
|
259
|
+
def _encode_sequence(seq, max_length=399):
|
|
260
|
+
vocab_map = {'A': 0, 'T': 1, 'C': 2, 'G': 3, 'N': 4}
|
|
261
|
+
encoded = [vocab_map.get(base, 4) for base in seq]
|
|
262
|
+
if len(encoded) > max_length:
|
|
263
|
+
encoded = encoded[:max_length]
|
|
264
|
+
else:
|
|
265
|
+
encoded += [4] * (max_length - len(encoded))
|
|
266
|
+
return np.array(encoded, dtype=np.float32).reshape(1, max_length, 1)
|
|
267
|
+
|
|
268
|
+
def _predict_long_torch(self, full_seq):
|
|
269
|
+
processed = PRFPredictor._process_sequence(full_seq)
|
|
270
|
+
encoded = PRFPredictor._encode_sequence(processed, max_length=self.long_seq_length)
|
|
271
|
+
x = torch.from_numpy(encoded).to(self.device)
|
|
272
|
+
with torch.no_grad():
|
|
273
|
+
prob = self.long_model(x).detach().cpu().numpy().reshape(-1)[0]
|
|
274
|
+
# 保证在 [0,1]
|
|
275
|
+
prob = float(np.clip(prob, 0.0, 1.0))
|
|
276
|
+
return prob
|
|
277
|
+
|
|
278
|
+
def predict_sequence(self, sequence, window_size=3, short_threshold=0.1, ensemble_weight=0.4):
|
|
279
|
+
"""
|
|
280
|
+
预测完整序列中的PRF位点(滑动窗口方法)
|
|
281
|
+
|
|
282
|
+
Args:
|
|
283
|
+
sequence: 输入DNA序列
|
|
284
|
+
window_size: 滑动窗口大小 (默认为3)
|
|
285
|
+
short_threshold: short模型概率阈值 (默认为0.1)
|
|
286
|
+
ensemble_weight: short模型在集成中的权重 (默认为0.4)
|
|
287
|
+
|
|
288
|
+
Returns:
|
|
289
|
+
pd.DataFrame: 包含预测结果的DataFrame
|
|
290
|
+
"""
|
|
291
|
+
if window_size < 1:
|
|
292
|
+
raise ValueError("窗口大小必须大于等于1")
|
|
293
|
+
if short_threshold < 0:
|
|
294
|
+
raise ValueError("short模型阈值必须大于等于0")
|
|
295
|
+
if not (0.0 <= ensemble_weight <= 1.0):
|
|
296
|
+
raise ValueError("ensemble_weight 必须在 0.0 到 1.0 之间")
|
|
297
|
+
|
|
298
|
+
results = []
|
|
299
|
+
long_weight = 1.0 - ensemble_weight
|
|
300
|
+
|
|
301
|
+
try:
|
|
302
|
+
# 确保序列为字符串并转换为大写
|
|
303
|
+
sequence = str(sequence).upper()
|
|
304
|
+
|
|
305
|
+
# 滑动窗口预测
|
|
306
|
+
for pos in range(0, len(sequence) - 2, window_size):
|
|
307
|
+
# 提取窗口序列
|
|
308
|
+
fs_period, full_seq = extract_window_sequences(sequence, pos)
|
|
309
|
+
|
|
310
|
+
if fs_period is None or full_seq is None:
|
|
311
|
+
continue
|
|
312
|
+
|
|
313
|
+
# 预测并记录结果
|
|
314
|
+
pred = self.predict_single_position(fs_period, full_seq, short_threshold, ensemble_weight)
|
|
315
|
+
pred.update({
|
|
316
|
+
'Position': pos,
|
|
317
|
+
'Codon': sequence[pos:pos+3],
|
|
318
|
+
'Short_Sequence': fs_period, # 更清晰的命名
|
|
319
|
+
'Long_Sequence': full_seq # 更清晰的命名
|
|
320
|
+
})
|
|
321
|
+
results.append(pred)
|
|
322
|
+
|
|
323
|
+
# 创建结果DataFrame
|
|
324
|
+
results_df = pd.DataFrame(results)
|
|
325
|
+
|
|
326
|
+
return results_df
|
|
327
|
+
|
|
328
|
+
except Exception as e:
|
|
329
|
+
raise Exception(f"序列预测过程出错: {str(e)}")
|
|
330
|
+
|
|
331
|
+
def plot_sequence_prediction(self, sequence, window_size=3, short_threshold=0.65,
|
|
332
|
+
long_threshold=0.8, ensemble_weight=0.4, title=None, save_path=None,
|
|
333
|
+
figsize=(12, 8), dpi=300):
|
|
334
|
+
"""
|
|
335
|
+
Plot sequence PRF prediction results
|
|
336
|
+
|
|
337
|
+
Args:
|
|
338
|
+
sequence: Input DNA sequence
|
|
339
|
+
window_size: Sliding window size (default: 3)
|
|
340
|
+
short_threshold: Short model (HistGB) filtering threshold (default: 0.65)
|
|
341
|
+
long_threshold: Long model (BiLSTM-CNN) filtering threshold (default: 0.8)
|
|
342
|
+
ensemble_weight: Weight of short model in ensemble (default: 0.4)
|
|
343
|
+
title: Plot title (optional)
|
|
344
|
+
save_path: Save path (optional, saves plot if provided)
|
|
345
|
+
figsize: Figure size (default: (12, 8))
|
|
346
|
+
dpi: Figure resolution (default: 300)
|
|
347
|
+
|
|
348
|
+
Returns:
|
|
349
|
+
tuple: (pd.DataFrame, matplotlib.figure.Figure) prediction results and figure object
|
|
350
|
+
"""
|
|
351
|
+
try:
|
|
352
|
+
# Validate weight parameter
|
|
353
|
+
if not (0.0 <= ensemble_weight <= 1.0):
|
|
354
|
+
raise ValueError("ensemble_weight must be between 0.0 and 1.0")
|
|
355
|
+
|
|
356
|
+
long_weight = 1.0 - ensemble_weight
|
|
357
|
+
|
|
358
|
+
# Get prediction results
|
|
359
|
+
results_df = self.predict_sequence(sequence, window_size=window_size,
|
|
360
|
+
short_threshold=0.1, ensemble_weight=ensemble_weight)
|
|
361
|
+
|
|
362
|
+
if results_df.empty:
|
|
363
|
+
raise ValueError("Prediction results are empty, please check input sequence")
|
|
364
|
+
|
|
365
|
+
# Get sequence length
|
|
366
|
+
seq_length = len(sequence)
|
|
367
|
+
|
|
368
|
+
# Calculate display width
|
|
369
|
+
desired_visual_width = max(3, seq_length // 100) # FS site width ~1% of sequence length
|
|
370
|
+
prob_width = max(1, desired_visual_width // 3) # Prediction probability width is 1/3 of FS site width
|
|
371
|
+
|
|
372
|
+
# Create figure with three subplots, set height ratios
|
|
373
|
+
fig = plt.figure(figsize=figsize)
|
|
374
|
+
|
|
375
|
+
# Set title
|
|
376
|
+
if title:
|
|
377
|
+
fig.suptitle(title, y=0.95, fontsize=10)
|
|
378
|
+
else:
|
|
379
|
+
fig.suptitle(f'PRF Prediction Results (Weights {ensemble_weight:.1f}:{long_weight:.1f})', y=0.95, fontsize=10)
|
|
380
|
+
|
|
381
|
+
# Adjust subplot ratios, make top two heatmaps smaller
|
|
382
|
+
gs = fig.add_gridspec(3, 1, height_ratios=[0.1, 0.1, 1], hspace=0.2)
|
|
383
|
+
|
|
384
|
+
# FS site heatmap - using fixed width, no blur effect
|
|
385
|
+
ax0 = fig.add_subplot(gs[0])
|
|
386
|
+
fs_data = np.zeros((1, seq_length))
|
|
387
|
+
# Note: No actual FS site information in sliding window prediction, so keep empty or show predicted sites
|
|
388
|
+
# Show high-confidence predictions as potential FS sites
|
|
389
|
+
for _, row in results_df.iterrows():
|
|
390
|
+
pos = int(row['Position'])
|
|
391
|
+
if (row['Short_Probability'] >= short_threshold and
|
|
392
|
+
row['Long_Probability'] >= long_threshold and
|
|
393
|
+
row['Ensemble_Probability'] >= 0.8): # High confidence threshold
|
|
394
|
+
half_width = desired_visual_width // 2
|
|
395
|
+
start_pos = max(0, pos - half_width)
|
|
396
|
+
end_pos = min(seq_length, pos + half_width + 1)
|
|
397
|
+
fs_data[0, start_pos:end_pos] = 1 # Use fixed value, no gradient
|
|
398
|
+
|
|
399
|
+
ax0.imshow(fs_data, cmap='Reds', aspect='auto', interpolation='nearest')
|
|
400
|
+
ax0.set_xticks([])
|
|
401
|
+
ax0.set_yticks([])
|
|
402
|
+
ax0.set_title('FS site', pad=2, fontsize=8)
|
|
403
|
+
|
|
404
|
+
# Prediction probability heatmap - using fixed width to display probabilities
|
|
405
|
+
ax1 = fig.add_subplot(gs[1])
|
|
406
|
+
prob_data = np.zeros((1, seq_length))
|
|
407
|
+
|
|
408
|
+
# Apply dual threshold filtering
|
|
409
|
+
for _, row in results_df.iterrows():
|
|
410
|
+
pos = int(row['Position'])
|
|
411
|
+
if (row['Short_Probability'] >= short_threshold and
|
|
412
|
+
row['Long_Probability'] >= long_threshold):
|
|
413
|
+
# Set fixed width for each probability value
|
|
414
|
+
start = max(0, pos - prob_width//2)
|
|
415
|
+
end = min(seq_length, pos + prob_width//2 + 1)
|
|
416
|
+
prob_data[0, start:end] = row['Ensemble_Probability']
|
|
417
|
+
|
|
418
|
+
im = ax1.imshow(prob_data, cmap='Reds', aspect='auto', vmin=0, vmax=1, interpolation='nearest')
|
|
419
|
+
ax1.set_xticks([])
|
|
420
|
+
ax1.set_yticks([])
|
|
421
|
+
ax1.set_title('Prediction', pad=2, fontsize=8)
|
|
422
|
+
|
|
423
|
+
# Main plot (bar chart)
|
|
424
|
+
ax2 = fig.add_subplot(gs[2])
|
|
425
|
+
|
|
426
|
+
# Apply filtering thresholds
|
|
427
|
+
filtered_probs = results_df['Ensemble_Probability'].copy()
|
|
428
|
+
mask = ((results_df['Short_Probability'] < short_threshold) |
|
|
429
|
+
(results_df['Long_Probability'] < long_threshold))
|
|
430
|
+
filtered_probs[mask] = 0
|
|
431
|
+
|
|
432
|
+
# Draw bar chart - use black color and alpha=0.6 to match prediction_sample style
|
|
433
|
+
ax2.bar(results_df['Position'], filtered_probs,
|
|
434
|
+
alpha=0.6, color='black', width=1.0)
|
|
435
|
+
|
|
436
|
+
# Set x-axis ticks
|
|
437
|
+
step = max(seq_length // 10, 50)
|
|
438
|
+
ax2.set_xticks(np.arange(0, seq_length, step))
|
|
439
|
+
ax2.tick_params(axis='x', rotation=45)
|
|
440
|
+
|
|
441
|
+
# Set labels
|
|
442
|
+
ax2.set_xlabel('Position')
|
|
443
|
+
ax2.set_ylabel('Probability')
|
|
444
|
+
|
|
445
|
+
# Set y-axis range
|
|
446
|
+
ax2.set_ylim(0, 1)
|
|
447
|
+
|
|
448
|
+
# Add grid
|
|
449
|
+
ax2.grid(True, alpha=0.3)
|
|
450
|
+
|
|
451
|
+
# Ensure all subplots have consistent x-axis range
|
|
452
|
+
for ax in [ax0, ax1, ax2]:
|
|
453
|
+
ax.set_xlim(-1, seq_length)
|
|
454
|
+
|
|
455
|
+
# Adjust layout
|
|
456
|
+
plt.tight_layout()
|
|
457
|
+
|
|
458
|
+
# Save plot if save path is provided
|
|
459
|
+
if save_path:
|
|
460
|
+
save_path = os.fspath(save_path)
|
|
461
|
+
plt.savefig(save_path, dpi=dpi, bbox_inches='tight')
|
|
462
|
+
# Also save PDF version
|
|
463
|
+
if save_path.endswith('.png'):
|
|
464
|
+
pdf_path = save_path.replace('.png', '.pdf')
|
|
465
|
+
plt.savefig(pdf_path, bbox_inches='tight')
|
|
466
|
+
print(f"Plot saved to: {save_path}")
|
|
467
|
+
|
|
468
|
+
return results_df, fig
|
|
469
|
+
|
|
470
|
+
except Exception as e:
|
|
471
|
+
raise Exception(f"Error plotting sequence prediction: {str(e)}")
|
|
472
|
+
|
|
473
|
+
def predict_regions(self, sequences, short_threshold=0.1, ensemble_weight=0.4):
|
|
474
|
+
'''
|
|
475
|
+
Predict region sequences (batch prediction of known 399bp sequences)
|
|
476
|
+
|
|
477
|
+
Args:
|
|
478
|
+
sequences: 399bp sequences or DataFrame/Series/list containing 399bp sequences
|
|
479
|
+
short_threshold: Short model probability threshold (default: 0.1)
|
|
480
|
+
ensemble_weight: Weight of short model in ensemble (default: 0.4)
|
|
481
|
+
|
|
482
|
+
Returns:
|
|
483
|
+
DataFrame: DataFrame containing prediction probabilities for all sequences
|
|
484
|
+
'''
|
|
485
|
+
try:
|
|
486
|
+
# Validate weight parameter
|
|
487
|
+
if not (0.0 <= ensemble_weight <= 1.0):
|
|
488
|
+
raise ValueError("ensemble_weight must be between 0.0 and 1.0")
|
|
489
|
+
|
|
490
|
+
# Unify input format
|
|
491
|
+
if isinstance(sequences, pd.DataFrame):
|
|
492
|
+
if 'Long_Sequence' in sequences.columns:
|
|
493
|
+
sequences = sequences['Long_Sequence']
|
|
494
|
+
elif '399bp' in sequences.columns:
|
|
495
|
+
sequences = sequences['399bp']
|
|
496
|
+
else:
|
|
497
|
+
raise ValueError("DataFrame must contain 'Long_Sequence' or '399bp' column")
|
|
498
|
+
if isinstance(sequences, pd.Series):
|
|
499
|
+
sequences = sequences.tolist()
|
|
500
|
+
elif isinstance(sequences, str):
|
|
501
|
+
sequences = [sequences]
|
|
502
|
+
|
|
503
|
+
results = []
|
|
504
|
+
for i, seq399 in enumerate(sequences):
|
|
505
|
+
try:
|
|
506
|
+
# Extract central 33bp from 399bp sequence (for short model use)
|
|
507
|
+
seq33 = self._extract_center_sequence(seq399, target_length=self.short_seq_length)
|
|
508
|
+
|
|
509
|
+
# Use unified prediction method
|
|
510
|
+
pred_result = self.predict_single_position(seq33, seq399, short_threshold, ensemble_weight)
|
|
511
|
+
pred_result.update({
|
|
512
|
+
'Short_Sequence': seq33,
|
|
513
|
+
'Long_Sequence': seq399
|
|
514
|
+
})
|
|
515
|
+
|
|
516
|
+
results.append(pred_result)
|
|
517
|
+
|
|
518
|
+
except Exception as e:
|
|
519
|
+
print(f"Error processing sequence {i+1}: {str(e)}")
|
|
520
|
+
long_weight = 1.0 - ensemble_weight
|
|
521
|
+
results.append({
|
|
522
|
+
'Short_Probability': 0.0,
|
|
523
|
+
'Long_Probability': 0.0,
|
|
524
|
+
'Ensemble_Probability': 0.0,
|
|
525
|
+
'Ensemble_Weights': f'Short:{ensemble_weight:.1f}, Long:{long_weight:.1f}',
|
|
526
|
+
'Short_Sequence': self._extract_center_sequence(seq399, target_length=self.short_seq_length) if len(seq399) >= self.short_seq_length else seq399,
|
|
527
|
+
'Long_Sequence': seq399
|
|
528
|
+
})
|
|
529
|
+
|
|
530
|
+
return pd.DataFrame(results)
|
|
531
|
+
|
|
532
|
+
except Exception as e:
|
|
533
|
+
raise Exception(f"Error in region prediction process: {str(e)}")
|
|
534
|
+
|
|
535
|
+
def _extract_center_sequence(self, sequence, target_length=33):
|
|
536
|
+
"""Extract subsequence of specified length from center position of sequence"""
|
|
537
|
+
# Ensure sequence is string
|
|
538
|
+
sequence = str(sequence).upper()
|
|
539
|
+
|
|
540
|
+
# If sequence length is less than target length, return original sequence
|
|
541
|
+
if len(sequence) <= target_length:
|
|
542
|
+
return sequence
|
|
543
|
+
|
|
544
|
+
# Calculate center position
|
|
545
|
+
center = len(sequence) // 2
|
|
546
|
+
half_target = target_length // 2
|
|
547
|
+
|
|
548
|
+
# Extract center sequence
|
|
549
|
+
start = center - half_target
|
|
550
|
+
end = start + target_length
|
|
551
|
+
|
|
552
|
+
# Boundary check
|
|
553
|
+
if start < 0:
|
|
554
|
+
start = 0
|
|
555
|
+
end = target_length
|
|
556
|
+
elif end > len(sequence):
|
|
557
|
+
end = len(sequence)
|
|
558
|
+
start = end - target_length
|
|
559
|
+
|
|
560
|
+
return sequence[start:end]
|
|
561
|
+
|
|
562
|
+
# 兼容性方法(向后兼容,但标记为废弃)
|
|
563
|
+
def predict_full(self, sequence, window_size=3, short_threshold=0.1, short_weight=0.4, plot=False):
|
|
564
|
+
"""
|
|
565
|
+
⚠️ 已废弃:请使用 predict_sequence() 方法
|
|
566
|
+
|
|
567
|
+
向后兼容的方法,内部调用新的 predict_sequence()
|
|
568
|
+
"""
|
|
569
|
+
import warnings
|
|
570
|
+
warnings.warn("predict_full() 已废弃,请使用 predict_sequence() 方法", DeprecationWarning, stacklevel=2)
|
|
571
|
+
|
|
572
|
+
# 调用新方法并添加兼容性字段
|
|
573
|
+
results_df = self.predict_sequence(sequence, window_size, short_threshold, short_weight)
|
|
574
|
+
|
|
575
|
+
# 添加兼容性字段
|
|
576
|
+
if 'Ensemble_Probability' in results_df.columns:
|
|
577
|
+
results_df['Voting_Probability'] = results_df['Ensemble_Probability']
|
|
578
|
+
results_df['Weighted_Probability'] = results_df['Ensemble_Probability']
|
|
579
|
+
if 'Ensemble_Weights' in results_df.columns:
|
|
580
|
+
results_df['Weight_Info'] = results_df['Ensemble_Weights']
|
|
581
|
+
if 'Short_Sequence' in results_df.columns:
|
|
582
|
+
results_df['33bp'] = results_df['Short_Sequence']
|
|
583
|
+
if 'Long_Sequence' in results_df.columns:
|
|
584
|
+
results_df['399bp'] = results_df['Long_Sequence']
|
|
585
|
+
|
|
586
|
+
if plot:
|
|
587
|
+
# 如果需要绘图,调用绘图方法
|
|
588
|
+
_, fig = self.plot_sequence_prediction(sequence, window_size, 0.65, 0.8, short_weight)
|
|
589
|
+
return results_df, fig
|
|
590
|
+
|
|
591
|
+
return results_df
|
|
592
|
+
|
|
593
|
+
def predict_region(self, seq, short_threshold=0.1, short_weight=0.4):
|
|
594
|
+
"""
|
|
595
|
+
⚠️ 已废弃:请使用 predict_regions() 方法
|
|
596
|
+
|
|
597
|
+
向后兼容的方法,内部调用新的 predict_regions()
|
|
598
|
+
"""
|
|
599
|
+
import warnings
|
|
600
|
+
warnings.warn("predict_region() 已废弃,请使用 predict_regions() 方法", DeprecationWarning, stacklevel=2)
|
|
601
|
+
|
|
602
|
+
# 调用新方法并添加兼容性字段
|
|
603
|
+
results_df = self.predict_regions(seq, short_threshold, short_weight)
|
|
604
|
+
|
|
605
|
+
# 添加兼容性字段
|
|
606
|
+
if 'Ensemble_Probability' in results_df.columns:
|
|
607
|
+
results_df['Voting_Probability'] = results_df['Ensemble_Probability']
|
|
608
|
+
results_df['Weighted_Probability'] = results_df['Ensemble_Probability']
|
|
609
|
+
if 'Ensemble_Weights' in results_df.columns:
|
|
610
|
+
results_df['Weight_Info'] = results_df['Ensemble_Weights']
|
|
611
|
+
if 'Short_Sequence' in results_df.columns:
|
|
612
|
+
results_df['33bp'] = results_df['Short_Sequence']
|
|
613
|
+
if 'Long_Sequence' in results_df.columns:
|
|
614
|
+
results_df['399bp'] = results_df['Long_Sequence']
|
|
615
|
+
|
|
616
|
+
return results_df
|
|
Binary file
|