Spaces:
Sleeping
Sleeping
| """Data loading and management for DynaCLR visualization.""" | |
| import pandas as pd | |
| import anndata as ad | |
| from iohub import open_ome_zarr | |
| from huggingface_hub import snapshot_download | |
| from pathlib import Path | |
| from .config import ( | |
| OME_ZARR_PATH, | |
| ANNDATA_PATH, | |
| INFECTION_ANNOTATIONS_PATH, | |
| HF_DATASET_REPO, | |
| USE_HF_DATASET, | |
| DATA_PATH, | |
| ) | |
| # Global state for cached data | |
| ome_dataset = None | |
| adata_demo = None | |
| embedding_info = None | |
| all_embeddings = None | |
| embedding_to_key_idx = None | |
| current_plot_data = None | |
| # Track selection caching | |
| cached_fov_choices = None # List of all FOV names | |
| cached_tracks_by_fov = None # Dict: {fov_name: [track_ids]} | |
| cached_annotated_tracks = ( | |
| None # Set of (track_id, fov_name) with infected/uninfected cells | |
| ) | |
| def download_dataset_from_hf(): | |
| """Download dataset files from HuggingFace dataset repository. | |
| This function downloads the entire dataset repository to a local cache | |
| and returns paths to the downloaded files. | |
| Returns: | |
| tuple: (ome_zarr_path, anndata_path, infection_csv_path) | |
| """ | |
| print(f"Downloading dataset from HuggingFace: {HF_DATASET_REPO}") | |
| # Get HF token from environment (required for private repos) | |
| # token = os.getenv("HF_TOKEN") | |
| # if not token: | |
| # raise ValueError( | |
| # "HF_TOKEN environment variable not set. " | |
| # "Please set it to access the private dataset repository." | |
| # ) | |
| # Download the entire dataset repository | |
| # This will cache files locally and reuse them on subsequent runs | |
| cache_dir = snapshot_download( | |
| repo_id=HF_DATASET_REPO, | |
| repo_type="dataset", | |
| # token=token, | |
| cache_dir=str(DATA_PATH.parent / ".hf_cache"), | |
| ) | |
| cache_path = Path(cache_dir) | |
| print(f"Dataset cached to: {cache_path}") | |
| # Return paths to the downloaded files | |
| ome_zarr_path = cache_path / "dataset.zarr" | |
| anndata_path = cache_path / "annotations_filtered.zarr" | |
| infection_csv_path = cache_path / "track_infection_annotation.csv" | |
| # Verify all files exist | |
| for path, name in [ | |
| (ome_zarr_path, "dataset.zarr"), | |
| (anndata_path, "annotations_filtered.zarr"), | |
| (infection_csv_path, "track_infection_annotation.csv"), | |
| ]: | |
| if not path.exists(): | |
| raise FileNotFoundError(f"Expected file not found in dataset: {name}") | |
| return ome_zarr_path, anndata_path, infection_csv_path | |
| def get_data_paths(): | |
| """Get data file paths, downloading from HF if USE_HF_DATASET is enabled. | |
| Returns: | |
| tuple: (ome_zarr_path, anndata_path, infection_csv_path) | |
| """ | |
| if USE_HF_DATASET: | |
| return download_dataset_from_hf() | |
| else: | |
| return OME_ZARR_PATH, ANNDATA_PATH, INFECTION_ANNOTATIONS_PATH | |
| def load_ome_dataset(): | |
| """Load OME-Zarr dataset once and reuse the handle.""" | |
| global ome_dataset | |
| if ome_dataset is None: | |
| ome_zarr_path, _, _ = get_data_paths() | |
| print(f"Loading OME-Zarr dataset from: {ome_zarr_path}") | |
| ome_dataset = open_ome_zarr(ome_zarr_path, mode="r") | |
| return ome_dataset | |
| def load_anndata(): | |
| """Load AnnData and extract embedding information with infection annotations.""" | |
| _, anndata_path, infection_csv_path = get_data_paths() | |
| print(f"Loading AnnData from: {anndata_path}") | |
| adata = ad.read_zarr(anndata_path) | |
| # Check if infection_status is already in obs (e.g., from filtered dataset) | |
| if "infection_status" not in adata.obs.columns: | |
| # Load infection annotations and merge with AnnData.obs | |
| print(f"Loading infection annotations from: {infection_csv_path}") | |
| infection_annotations = pd.read_csv(infection_csv_path) | |
| # Merge infection status using (fov_name, id) as the composite unique key | |
| # This is necessary because 'id' is only unique within each FOV | |
| print("Merging infection annotations with AnnData.obs...") | |
| adata.obs = adata.obs.merge( | |
| infection_annotations[["fov_name", "id", "infection_status"]], | |
| on=["fov_name", "id"], | |
| how="left", | |
| ) | |
| # Fill missing infection_status with 'unknown' | |
| adata.obs["infection_status"] = adata.obs["infection_status"].fillna("unknown") | |
| # Report infection status statistics | |
| n_with_annot = (adata.obs["infection_status"] != "unknown").sum() | |
| n_without_annot = (adata.obs["infection_status"] == "unknown").sum() | |
| print( | |
| f" - Cells with infection annotations: {n_with_annot} " | |
| f"({n_with_annot / len(adata.obs) * 100:.1f}%)" | |
| ) | |
| print( | |
| f" - Cells without annotations: {n_without_annot} " | |
| f"({n_without_annot / len(adata.obs) * 100:.1f}%)" | |
| ) | |
| # Detect all embeddings from obsm | |
| embeddings = {} | |
| for key in adata.obsm.keys(): | |
| n_components = adata.obsm[key].shape[1] | |
| if key == "X_pca": | |
| # PCA components: PC1, PC2, ..., PC8 | |
| embeddings[key] = [f"PC{i + 1}" for i in range(n_components)] | |
| elif key == "X_projections": | |
| # Projection dimensions: Proj1, Proj2, ..., Proj32 | |
| embeddings[key] = [f"Proj{i + 1}" for i in range(n_components)] | |
| else: | |
| # Generic naming for any other embeddings | |
| embeddings[key] = [f"{key}_{i + 1}" for i in range(n_components)] | |
| print(f"Loaded {adata.shape[0]} cells with {adata.shape[1]} features") | |
| print(f"Available embeddings: {list(embeddings.keys())}") | |
| return adata, embeddings | |
| def initialize_data(): | |
| """Initialize all data and create embedding index.""" | |
| global adata_demo, embedding_info, all_embeddings, embedding_to_key_idx | |
| global cached_fov_choices, cached_tracks_by_fov, cached_annotated_tracks | |
| adata_demo, embedding_info = load_anndata() | |
| # Create flat list of all embedding names for dropdowns | |
| all_embeddings = [] | |
| embedding_to_key_idx = {} | |
| for obsm_key, component_names in embedding_info.items(): | |
| for idx, component_name in enumerate(component_names): | |
| all_embeddings.append(component_name) | |
| embedding_to_key_idx[component_name] = (obsm_key, idx) | |
| print(f"Total embedding dimensions available: {len(all_embeddings)}") | |
| # Pre-compute FOV and track selections with filtering | |
| print("Pre-computing FOV and track selections...") | |
| # Filter tracks: keep only those with at least one infected or uninfected cell | |
| print("Filtering tracks with infection annotations...") | |
| track_groups = adata_demo.obs.groupby(["track_id", "fov_name"])["infection_status"] | |
| cached_annotated_tracks = set() | |
| for (track_id, fov_name), statuses in track_groups: | |
| # Include track if it has at least one infected or uninfected cell | |
| if any(status in ["infected", "uninfected"] for status in statuses): | |
| cached_annotated_tracks.add((int(track_id), fov_name)) | |
| total_tracks = len(track_groups) | |
| annotated_count = len(cached_annotated_tracks) | |
| print( | |
| f"Filtered tracks: {annotated_count} / {total_tracks} have infection annotations " | |
| f"({annotated_count / total_tracks * 100:.1f}%)" | |
| ) | |
| # Build FOV → Tracks mapping (only annotated tracks) | |
| cached_tracks_by_fov = {} | |
| for track_id, fov_name in cached_annotated_tracks: | |
| if fov_name not in cached_tracks_by_fov: | |
| cached_tracks_by_fov[fov_name] = [] | |
| cached_tracks_by_fov[fov_name].append(track_id) | |
| # Sort track lists within each FOV | |
| for fov_name in cached_tracks_by_fov: | |
| cached_tracks_by_fov[fov_name].sort() | |
| # Only show FOVs that have annotated tracks | |
| cached_fov_choices = sorted(cached_tracks_by_fov.keys()) | |
| total_fovs = len(adata_demo.obs["fov_name"].unique()) | |
| print( | |
| f"Cached annotated tracks for {len(cached_tracks_by_fov)} FOVs (out of {total_fovs} total FOVs)" | |
| ) | |
| return adata_demo, embedding_info, all_embeddings | |
| def get_embedding_data(embedding_name): | |
| """Extract embedding data for a given component name.""" | |
| obsm_key, component_idx = embedding_to_key_idx[embedding_name] | |
| return adata_demo.obsm[obsm_key][:, component_idx] | |
| def get_all_embeddings(): | |
| """Get list of all available embedding names.""" | |
| return all_embeddings | |
| def get_fov_choices(): | |
| """Get pre-computed list of FOV choices for dropdown.""" | |
| return cached_fov_choices or [] | |
| def get_tracks_for_fov(fov_name): | |
| """Get list of annotated track IDs for a specific FOV. | |
| Args: | |
| fov_name: FOV name (e.g., "A/1/000000") | |
| Returns: | |
| list: Sorted list of track IDs with infection annotations in this FOV | |
| """ | |
| if cached_tracks_by_fov is None: | |
| return [] | |
| return cached_tracks_by_fov.get(fov_name, []) | |
| def is_track_annotated(track_id, fov_name): | |
| """Check if a track has infection annotations. | |
| Args: | |
| track_id: Track ID | |
| fov_name: FOV name | |
| Returns: | |
| bool: True if track has at least one infected or uninfected cell | |
| """ | |
| if cached_annotated_tracks is None: | |
| return False | |
| return (int(track_id), fov_name) in cached_annotated_tracks | |