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/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
@@ -0,0 +1,4 @@
1
+ [diffend] Oversized file quarantined before diffing.
2
+ name: FScanpy/pretrained/long.pth
3
+ size: 39531146 bytes
4
+ sha256: 1207fab6a80c9f8f887b3bff2062d9eb20df96387d70207142868fbe65c4fd4a
Binary file