Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 16 additions & 1 deletion examples/example.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,7 +26,7 @@

# subset
project_meta_dataframe = recount_meta_dataframe.filter(
pl.col("project").is_in(["SRP009615"]) & pl.col("external_id").is_in(["SRR389077"])
pl.col("project").is_in(["SRP009615"])
)

print(project_meta_dataframe)
Expand Down Expand Up @@ -61,6 +61,21 @@
print(gene_annotation)
print(gene_counts)

scaled_counts = project.scale_mapped_reads(
gene_counts,
target_size=4e7,
L=100,
)

print(scaled_counts)

scaled_counts = project.scale_auc(
gene_counts,
target_size=4e7,
)

print(scaled_counts)

exon_annotation, exon_counts = project.load(Dtype.EXON)

print(exon_annotation)
Expand Down
74 changes: 67 additions & 7 deletions src/pyrecount/accessor.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,9 @@ class Project:

project_ids: List[str] = field(init=False)
sample: List[str] = field(init=False)
endpoints: "EndpointConnector" = field(init=False)
endpoints: EndpointConnector = field(init=False)

_cached_metadata: Optional[pl.DataFrame] = field(init=False, default=None)

def __post_init__(self):
if not isinstance(self.metadata, pl.DataFrame):
Expand Down Expand Up @@ -86,7 +88,6 @@ def get_project_urls(self, dtype) -> List[str]:
dtype=dtype,
annotation=self.annotation,
project_ids=self.project_ids,
# only needed for bigwig
sample=self.sample,
jxn_format=self.jxn_format,
)
Expand All @@ -108,13 +109,71 @@ async def cache(self) -> None:
makedirs(path.dirname(fpath), exist_ok=True)
tasks.append(download_url_to_path(url=url, fpath=fpath))

# launches all tasks, without rate limiting
if tasks:
await asyncio.gather(*tasks)

def scale_mapped_reads(self, counts: pl.DataFrame, target_size: float, L: int):
md = self.load_metadata()

mapped_reads = pl.col("star.all_mapped_reads").cast(pl.Float64)
avg_mapped_len = pl.col("star.average_mapped_length").cast(pl.Float64)
avg_read_len = pl.col("avg_len").cast(pl.Float64)

# paired-end detection
ratio = (avg_mapped_len / avg_read_len).round(0)
paired_end = ratio == 2
paired_factor = pl.when(paired_end).then(2).otherwise(1)

# scale factors
sf = md.select(
[
pl.col("external_id"),
(
target_size * L * paired_factor / (mapped_reads * avg_mapped_len**2)
).alias("sf"),
]
)

sf_map = dict(zip(sf["external_id"], sf["sf"]))

return counts.with_columns(
[
(pl.col(c) * sf_map[c])
for c in counts.select(pl.selectors.numeric()).columns
]
)

def scale_auc(self, counts: pl.DataFrame, target_size: float) -> pl.DataFrame:
md = self.load_metadata()

auc = pl.col("bc_auc.all_reads_all_bases").cast(pl.Float64)

# scale factors
sf = md.select(
"external_id",
(target_size / auc).alias("sf"),
)

sf_map = dict(zip(sf["external_id"], sf["sf"]))

return counts.with_columns(
[
pl.col(c).mul(sf_map[c]).round(0).cast(pl.Int64).alias(c)
for c in counts.columns
if c != "gene_id"
]
)

def load_metadata(self) -> pl.DataFrame:
if self._cached_metadata is None:
self._cached_metadata = self._metadata_load()
return self._cached_metadata

def load(self, dtype) -> Union[pl.DataFrame, Tuple[pl.DataFrame, pl.DataFrame]]:
match dtype:
case Dtype.METADATA:
return self._metadata_load()
return self.load_metadata()
case Dtype.JXN:
return self._jxn_load()
case Dtype.GENE:
Expand Down Expand Up @@ -315,9 +374,11 @@ def _gtf_read(self, rpath: str) -> pl.DataFrame:
[
annotation_dataframe["attribute"]
.map_elements(
lambda x: re.findall(rf'{field} "([^"]*)"', x)[0]
if rf'{field} "' in x
else "",
lambda x: (
re.findall(rf'{field} "([^"]*)"', x)[0]
if rf'{field} "' in x
else ""
),
return_dtype=pl.Utf8,
)
.alias(field)
Expand Down Expand Up @@ -366,7 +427,6 @@ def _exon_load(self) -> pl.DataFrame:
annotation = self._gtf_read(fpath)
if url.endswith(f"{self.annotation.value}.gz"):
counts = self._counts_read(fpath)
# TODO: extract first column (chromosome|start_1base|end_1ba…)
exon_colname = counts.columns[0]
exon_fields = ["chrom", "start", "end", "strand"]
counts = (
Expand Down
7 changes: 3 additions & 4 deletions tests/test_accessor.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,6 @@
from pyrecount.accessor import Metadata, Project
from pyrecount.models import Dtype, Annotation

# TODO: transform raw counts
# TODO: jxn dataframe headers, external_id not rail id
# TODO: multi-project support for exon, gene dtypes
# TODO: expand sra attributes
Expand Down Expand Up @@ -241,20 +240,20 @@ async def test_project_gene_accessor(
(pl.col("project").is_in(project_ids))
)

dtype = Dtype.GENE
dtype = [Dtype.METADATA, Dtype.GENE]

gene = Project(
metadata=project_meta_dataframe,
dbase=dbase,
organism=organism,
dtype=[dtype],
dtype=dtype,
annotation=annotation,
jxn_format=None,
root_url=root_url,
)

await gene.cache()
gene_annotation, gene_counts = gene.load(dtype)
gene_annotation, gene_counts = gene.load(Dtype.GENE)

assert gene_annotation.shape == expected_annotation_shape
assert gene_counts.shape == expected_counts_shape
Expand Down