| from dataclasses import dataclass |
| from typing import Any, Dict, Optional |
|
|
| import torch |
|
|
|
|
| @dataclass |
| class PopulationPerturbationBatch: |
| """Batched population-level perturbation data. |
| |
| Shapes |
| ------ |
| source_cells : [B, Ns, G] |
| target_cells : [B, Nt, G] |
| perturbation : [B, G] — multi-hot gene target indicator |
| source_mask : [B, Ns] — 1 = real cell, 0 = padding |
| target_mask : [B, Nt] — 1 = real cell, 0 = padding |
| """ |
|
|
| source_cells : torch.Tensor |
| target_cells : torch.Tensor |
| perturbation : torch.Tensor |
| source_mask : torch.Tensor |
| target_mask : torch.Tensor |
| context : Optional[Dict[str, Any]] = None |
| metadata : Optional[Dict[str, Any]] = None |
| |
| |
| cell_line_id : Optional[int] = None |
| |
| |
| prior_score : Optional[torch.Tensor] = None |
|
|
| def to(self, device: torch.device) -> "PopulationPerturbationBatch": |
| return PopulationPerturbationBatch( |
| source_cells=self.source_cells.to(device), |
| target_cells=self.target_cells.to(device), |
| perturbation=self.perturbation.to(device), |
| source_mask=self.source_mask.to(device), |
| target_mask=self.target_mask.to(device), |
| context=self.context, |
| metadata=self.metadata, |
| cell_line_id=self.cell_line_id, |
| prior_score=self.prior_score.to(device) if self.prior_score is not None else None, |
| ) |
|
|