calm-data-generator 2.2.0__tar.gz → 2.2.1__tar.gz
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.
- {calm_data_generator-2.2.0/calm_data_generator.egg-info → calm_data_generator-2.2.1}/PKG-INFO +1 -1
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/__init__.py +5 -5
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/RealBlockGenerator.py +8 -7
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/RealGenerator.py +256 -7
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/ExternalReporter.py +4 -2
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/QualityReporter.py +46 -39
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1/calm_data_generator.egg-info}/PKG-INFO +1 -1
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/pyproject.toml +1 -1
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_accessors.py +98 -21
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/LICENSE +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/MANIFEST.in +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/README.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/cli.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/API.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/API_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CAUSAL_ENGINE_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CAUSAL_ENGINE_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_BLOCK_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/COMPLEX_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/COMPLEX_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DOCUMENTATION.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DOCUMENTATION_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DRIFT_INJECTOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DRIFT_INJECTOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/PRESETS_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/PRESETS_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_BLOCK_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REPORTS_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REPORTS_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/SCENARIO_INJECTOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/SCENARIO_INJECTOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_BLOCK_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_GENERATOR_REFERENCE.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_GENERATOR_REFERENCE_ES.md +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/base.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/Clinic.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/ClinicGeneratorBlock.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/ClinicReporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/complex/ComplexGenerator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/complex/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/configs.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/drift/DriftInjector.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/drift/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/CausalEngine.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/ScenarioInjector.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/persistence_models.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/GeneratorFactory.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamBlockGenerator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamGenerator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamReporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/CustomPluginAdapter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/QualityReporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/utils/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/utils/propagation.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/logger.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/BalancePreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ConceptDriftPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/CopulaPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DataQualityAuditPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DiffusionPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DriftScenarioPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/FastPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/FastPrototypePreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/GradualDriftPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/HighFidelityPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ImbalancePreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/LongitudinalHealthPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/OmicsIntegrationPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/RareDiseasePreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ScenarioInjectorPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/SeasonalTimeSeriesPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/SingleCellQualityPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/TimeSeriesPreset.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/base.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/DiscriminatorReporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/LocalIndexGenerator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/Visualizer.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/__init__.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/base.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/advanced_drifts.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinic_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinical_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinical_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/correlation_drift.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/drift_injector.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/real_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/real_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/reports_deep_dive.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/scenario_injector.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/stream_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/stream_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/tutorial_advanced_methods.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/SOURCES.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/dependency_links.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/entry_points.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/requires.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/top_level.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/requirements.txt +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/setup.cfg +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_anndata_support.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_block_generators_extended.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_causal_engine.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinic_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinical_advanced.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinical_regression.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_complex_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_comprehensive.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_comprehensive_reporting.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_correlation_propagation.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_custom_plugin_adapter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_differentiation_factor.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_discriminator_reporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_disease_effects_fix.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_correlations.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_injector_advanced.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_injector_math.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_functional_drift.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_imbalance.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_migration.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_presets.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_quality_metrics.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_real_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_real_generator_persistence.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_reporters_extended.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_reporting_fix.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_river_integration.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scenario_extended.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scgft_reporter.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scvi_quality_regression.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_single_call.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_stream_block_generator.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_tabular_extended.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_timeseries_extended.py +0 -0
- {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_timeseries_real.py +0 -0
{calm_data_generator-2.2.0/calm_data_generator.egg-info → calm_data_generator-2.2.1}/PKG-INFO
RENAMED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: calm-data-generator
|
|
3
|
-
Version: 2.2.
|
|
3
|
+
Version: 2.2.1
|
|
4
4
|
Summary: CALM-Data-Generator: A Python library for synthetic data generation with support for drift injection and clinical data.
|
|
5
5
|
Author-email: Alejandro Belda Fernandez <alejandrobeldafernandez@gmail.com>
|
|
6
6
|
License: MIT
|
|
@@ -1,11 +1,11 @@
|
|
|
1
1
|
# CALM-Data-Generator - Synthetic Data Generation Library
|
|
2
2
|
|
|
3
|
-
from calm_data_generator
|
|
3
|
+
from calm_data_generator import presets
|
|
4
4
|
from calm_data_generator.generators.clinical import ClinicalDataGenerator
|
|
5
|
-
from calm_data_generator.generators.drift import DriftInjector
|
|
6
|
-
from calm_data_generator.generators.dynamics import ScenarioInjector, CausalEngine
|
|
7
5
|
from calm_data_generator.generators.complex import ComplexGenerator
|
|
8
|
-
from calm_data_generator import
|
|
6
|
+
from calm_data_generator.generators.drift import DriftInjector
|
|
7
|
+
from calm_data_generator.generators.dynamics import CausalEngine, ScenarioInjector
|
|
8
|
+
from calm_data_generator.generators.tabular import QualityReporter, RealGenerator
|
|
9
9
|
|
|
10
10
|
# Optional imports that may fail
|
|
11
11
|
try:
|
|
@@ -13,7 +13,7 @@ try:
|
|
|
13
13
|
except ImportError:
|
|
14
14
|
StreamGenerator = None
|
|
15
15
|
|
|
16
|
-
__version__ = "2.1
|
|
16
|
+
__version__ = "2.2.1"
|
|
17
17
|
|
|
18
18
|
__all__ = [
|
|
19
19
|
# Generators
|
|
@@ -13,11 +13,11 @@ Key Features:
|
|
|
13
13
|
- **Timestamp Alignment**: Can inject timestamps that are aligned with the block structure.
|
|
14
14
|
"""
|
|
15
15
|
|
|
16
|
-
import os
|
|
17
16
|
import logging
|
|
17
|
+
import os
|
|
18
18
|
import warnings
|
|
19
19
|
from pathlib import Path
|
|
20
|
-
from typing import
|
|
20
|
+
from typing import Any, Dict, List, Optional, Union
|
|
21
21
|
|
|
22
22
|
import numpy as np
|
|
23
23
|
import pandas as pd
|
|
@@ -27,9 +27,10 @@ warnings.filterwarnings("ignore", category=UserWarning)
|
|
|
27
27
|
warnings.filterwarnings("ignore", category=FutureWarning)
|
|
28
28
|
warnings.filterwarnings("ignore", category=DeprecationWarning)
|
|
29
29
|
|
|
30
|
-
from .
|
|
31
|
-
from calm_data_generator.reports.QualityReporter import QualityReporter
|
|
32
|
-
|
|
30
|
+
from calm_data_generator.generators.configs import DriftConfig, ReportConfig # noqa: E402
|
|
31
|
+
from calm_data_generator.reports.QualityReporter import QualityReporter # noqa: E402
|
|
32
|
+
|
|
33
|
+
from .RealGenerator import RealGenerator # noqa: E402
|
|
33
34
|
|
|
34
35
|
|
|
35
36
|
class RealBlockGenerator(RealGenerator):
|
|
@@ -59,7 +60,7 @@ class RealBlockGenerator(RealGenerator):
|
|
|
59
60
|
self.logger = logging.getLogger(self.__class__.__name__)
|
|
60
61
|
self.logger.setLevel(logging.INFO)
|
|
61
62
|
|
|
62
|
-
|
|
63
|
+
|
|
63
64
|
self.logger.info("RealBlockGenerator initialized.")
|
|
64
65
|
|
|
65
66
|
# --------------------------------------------------------------------------- #
|
|
@@ -162,10 +163,10 @@ class RealBlockGenerator(RealGenerator):
|
|
|
162
163
|
data=block_data_no_block,
|
|
163
164
|
method=method,
|
|
164
165
|
target_col=target_col,
|
|
165
|
-
model_params=model_params,
|
|
166
166
|
n_samples=n_samples,
|
|
167
167
|
output_dir=output_dir,
|
|
168
168
|
custom_distributions=custom_distributions,
|
|
169
|
+
**(model_params or {}),
|
|
169
170
|
)
|
|
170
171
|
|
|
171
172
|
if synthetic_block is None:
|
|
@@ -642,6 +642,250 @@ class RealGenerator(BaseGenerator):
|
|
|
642
642
|
|
|
643
643
|
return None
|
|
644
644
|
|
|
645
|
+
def encode_to_latent(
|
|
646
|
+
self,
|
|
647
|
+
data,
|
|
648
|
+
target_col: Optional[str] = None,
|
|
649
|
+
) -> "torch.Tensor":
|
|
650
|
+
"""
|
|
651
|
+
Encodes ``data`` into the model's latent space, handling the full
|
|
652
|
+
preprocessing pipeline for each supported method.
|
|
653
|
+
|
|
654
|
+
- **tvae / rtvae**: runs TabularEncoder → VAE encoder → returns ``mu``
|
|
655
|
+
(mean of the approximate posterior). The TabularEncoder was fitted
|
|
656
|
+
on the *complete* training DataFrame (including ``target_col``), so
|
|
657
|
+
``data`` must contain the same columns, in the same order, as the
|
|
658
|
+
training data — ``target_col`` included.
|
|
659
|
+
- **scvi / scanvi**: calls ``model.get_latent_representation()`` and
|
|
660
|
+
returns the result as a float32 tensor.
|
|
661
|
+
|
|
662
|
+
Unlike :meth:`get_encoder`, which exposes the raw PyTorch module, this
|
|
663
|
+
method handles the method-specific preprocessing automatically. Use it
|
|
664
|
+
when implementing external drift analyses that need latent
|
|
665
|
+
representations for both TVAE and SCVI without rewriting the encoding
|
|
666
|
+
pipeline each time.
|
|
667
|
+
|
|
668
|
+
Parameters
|
|
669
|
+
----------
|
|
670
|
+
data:
|
|
671
|
+
- TVAE/RTVAE: a :class:`pandas.DataFrame` with **all** columns
|
|
672
|
+
used during training (including ``target_col``).
|
|
673
|
+
- SCVI/SCANVI: a :class:`pandas.DataFrame` *or* an
|
|
674
|
+
``AnnData`` object whose ``var_names`` match the training genes.
|
|
675
|
+
target_col:
|
|
676
|
+
Label column used to build the optional conditioning tensor for
|
|
677
|
+
TVAE models trained with a conditional dimension. For SCVI this
|
|
678
|
+
argument is ignored.
|
|
679
|
+
|
|
680
|
+
Returns
|
|
681
|
+
-------
|
|
682
|
+
torch.Tensor
|
|
683
|
+
Float32 tensor of shape ``(n_samples, n_latent)``.
|
|
684
|
+
|
|
685
|
+
Raises
|
|
686
|
+
------
|
|
687
|
+
RuntimeError
|
|
688
|
+
If no model has been trained yet, or the method is not supported.
|
|
689
|
+
"""
|
|
690
|
+
if not self.synthesizer:
|
|
691
|
+
raise RuntimeError("No trained model found. Call generate() first.")
|
|
692
|
+
|
|
693
|
+
if self.method in ("tvae", "rtvae"):
|
|
694
|
+
tabular_model = self.synthesizer.model
|
|
695
|
+
pytorch_model = tabular_model.model
|
|
696
|
+
pytorch_model.eval()
|
|
697
|
+
|
|
698
|
+
# TVAE trains the TabularEncoder on the FULL DataFrame (including
|
|
699
|
+
# target_col). We must pass the same columns here; do not drop
|
|
700
|
+
# target_col before encoding.
|
|
701
|
+
data_encoded = tabular_model.encode(data)
|
|
702
|
+
data_tensor = torch.tensor(
|
|
703
|
+
data_encoded.values, dtype=torch.float32
|
|
704
|
+
).to(pytorch_model.device)
|
|
705
|
+
|
|
706
|
+
# Append conditional dimensions expected by the encoder.
|
|
707
|
+
cond_dim = getattr(pytorch_model, "n_units_conditional", 0)
|
|
708
|
+
if cond_dim > 0 and target_col and isinstance(data, pd.DataFrame) and target_col in data.columns:
|
|
709
|
+
unique_classes = data[target_col].unique()
|
|
710
|
+
label_to_idx = {str(c): i for i, c in enumerate(unique_classes)}
|
|
711
|
+
label_indices = torch.tensor(
|
|
712
|
+
[label_to_idx.get(str(l), 0) for l in data[target_col].values],
|
|
713
|
+
dtype=torch.long,
|
|
714
|
+
device=pytorch_model.device,
|
|
715
|
+
)
|
|
716
|
+
cond_tensor = torch.zeros(
|
|
717
|
+
len(data), cond_dim, device=pytorch_model.device
|
|
718
|
+
)
|
|
719
|
+
cond_tensor.scatter_(
|
|
720
|
+
1, label_indices.unsqueeze(1).clamp(max=cond_dim - 1), 1.0
|
|
721
|
+
)
|
|
722
|
+
data_tensor = torch.cat([data_tensor, cond_tensor], dim=1)
|
|
723
|
+
|
|
724
|
+
with torch.no_grad():
|
|
725
|
+
encoder_out = pytorch_model.encoder(data_tensor)
|
|
726
|
+
mu = encoder_out[0] if isinstance(encoder_out, (tuple, list)) else encoder_out
|
|
727
|
+
return mu
|
|
728
|
+
|
|
729
|
+
if self.method in ("scvi", "scanvi"):
|
|
730
|
+
import torch as _torch
|
|
731
|
+
model = self.synthesizer
|
|
732
|
+
z = model.get_latent_representation()
|
|
733
|
+
return _torch.tensor(z, dtype=_torch.float32)
|
|
734
|
+
|
|
735
|
+
raise RuntimeError(
|
|
736
|
+
f"encode_to_latent() is not supported for method '{self.method}'. "
|
|
737
|
+
"Supported: tvae, rtvae, scvi, scanvi."
|
|
738
|
+
)
|
|
739
|
+
|
|
740
|
+
def decode_from_latent(
|
|
741
|
+
self,
|
|
742
|
+
z: "torch.Tensor",
|
|
743
|
+
data=None,
|
|
744
|
+
target_col: Optional[str] = None,
|
|
745
|
+
) -> pd.DataFrame:
|
|
746
|
+
"""
|
|
747
|
+
Decodes latent vectors ``z`` back to the original feature space.
|
|
748
|
+
|
|
749
|
+
- **tvae / rtvae**: runs VAE decoder → TabularEncoder inverse_transform.
|
|
750
|
+
- **scvi / scanvi**: calls the model's generative pass and returns
|
|
751
|
+
a DataFrame with the original gene columns.
|
|
752
|
+
|
|
753
|
+
Parameters
|
|
754
|
+
----------
|
|
755
|
+
z:
|
|
756
|
+
Float32 tensor of shape ``(n_samples, n_latent)`` produced by
|
|
757
|
+
:meth:`encode_to_latent` (or a manually perturbed version of it).
|
|
758
|
+
data:
|
|
759
|
+
Reference data used to infer library size (SCVI only) and column
|
|
760
|
+
names. Can be a :class:`pandas.DataFrame` or ``AnnData``.
|
|
761
|
+
target_col:
|
|
762
|
+
Label column to attach to the output DataFrame. When provided
|
|
763
|
+
together with ``data``, labels are copied from the reference.
|
|
764
|
+
|
|
765
|
+
Returns
|
|
766
|
+
-------
|
|
767
|
+
pandas.DataFrame
|
|
768
|
+
Decoded samples in the original feature space.
|
|
769
|
+
|
|
770
|
+
Raises
|
|
771
|
+
------
|
|
772
|
+
RuntimeError
|
|
773
|
+
If no model has been trained yet, or the method is not supported.
|
|
774
|
+
"""
|
|
775
|
+
if not self.synthesizer:
|
|
776
|
+
raise RuntimeError("No trained model found. Call generate() first.")
|
|
777
|
+
|
|
778
|
+
n_samples = z.shape[0]
|
|
779
|
+
|
|
780
|
+
if self.method in ("tvae", "rtvae"):
|
|
781
|
+
tabular_model = self.synthesizer.model
|
|
782
|
+
pytorch_model = tabular_model.model
|
|
783
|
+
pytorch_model.eval()
|
|
784
|
+
|
|
785
|
+
cond_dim = getattr(pytorch_model, "n_units_conditional", 0)
|
|
786
|
+
cond_tensor = None
|
|
787
|
+
if cond_dim > 0 and target_col is not None and data is not None and isinstance(data, pd.DataFrame):
|
|
788
|
+
unique_classes = data[target_col].unique()
|
|
789
|
+
label_to_idx = {str(c): i for i, c in enumerate(unique_classes)}
|
|
790
|
+
label_indices = torch.tensor(
|
|
791
|
+
[label_to_idx.get(str(l), 0) for l in data[target_col].values],
|
|
792
|
+
dtype=torch.long,
|
|
793
|
+
device=pytorch_model.device,
|
|
794
|
+
)
|
|
795
|
+
cond_tensor = torch.zeros(
|
|
796
|
+
n_samples, cond_dim, device=pytorch_model.device
|
|
797
|
+
)
|
|
798
|
+
cond_tensor.scatter_(
|
|
799
|
+
1, label_indices.unsqueeze(1).clamp(max=cond_dim - 1), 1.0
|
|
800
|
+
)
|
|
801
|
+
|
|
802
|
+
z = z.to(pytorch_model.device)
|
|
803
|
+
with torch.no_grad():
|
|
804
|
+
reconstructed = pytorch_model.decoder(z, cond_tensor)
|
|
805
|
+
|
|
806
|
+
# Get encoded column names: encode the full data (TabularEncoder
|
|
807
|
+
# was fitted on the complete DataFrame including target_col).
|
|
808
|
+
enc_cols = tabular_model.encode(data).columns if data is not None else None
|
|
809
|
+
reconstructed_df = pd.DataFrame(
|
|
810
|
+
reconstructed.cpu().numpy(),
|
|
811
|
+
columns=enc_cols,
|
|
812
|
+
)
|
|
813
|
+
synth_df = tabular_model.decode(reconstructed_df)
|
|
814
|
+
return synth_df
|
|
815
|
+
|
|
816
|
+
if self.method in ("scvi", "scanvi"):
|
|
817
|
+
model = self.synthesizer
|
|
818
|
+
device = model.device if hasattr(model, "device") else next(model.module.parameters()).device
|
|
819
|
+
|
|
820
|
+
if data is not None and hasattr(data, "X"):
|
|
821
|
+
raw = data.X.toarray() if hasattr(data.X, "toarray") else np.array(data.X)
|
|
822
|
+
orig_log_lib = np.log(raw.sum(axis=1) + 1e-8)
|
|
823
|
+
elif data is not None and isinstance(data, pd.DataFrame):
|
|
824
|
+
num_cols = data.select_dtypes(include=[np.number]).columns
|
|
825
|
+
if target_col:
|
|
826
|
+
num_cols = [c for c in num_cols if c != target_col]
|
|
827
|
+
orig_log_lib = np.log(data[num_cols].values.sum(axis=1) + 1e-8)
|
|
828
|
+
else:
|
|
829
|
+
orig_log_lib = np.zeros(n_samples)
|
|
830
|
+
|
|
831
|
+
z = z.to(device)
|
|
832
|
+
library_tensor = torch.tensor(
|
|
833
|
+
orig_log_lib, dtype=torch.float32
|
|
834
|
+
).unsqueeze(1).to(device)
|
|
835
|
+
batch_index = torch.zeros(n_samples, 1, dtype=torch.long).to(device)
|
|
836
|
+
|
|
837
|
+
y_tensor = None
|
|
838
|
+
if getattr(model.module, "dispersion", "gene") == "gene-label" and data is not None:
|
|
839
|
+
try:
|
|
840
|
+
label_registry = model.adata_manager.get_state_registry("labels")
|
|
841
|
+
cat_mapping = label_registry.categorical_mapping
|
|
842
|
+
label_map = {str(cat): i for i, cat in enumerate(cat_mapping)}
|
|
843
|
+
if hasattr(data, "obs") and target_col in data.obs.columns:
|
|
844
|
+
labels = data.obs[target_col].astype(str).values
|
|
845
|
+
elif isinstance(data, pd.DataFrame) and target_col in data.columns:
|
|
846
|
+
labels = data[target_col].astype(str).values
|
|
847
|
+
else:
|
|
848
|
+
labels = None
|
|
849
|
+
if labels is not None:
|
|
850
|
+
y_tensor = torch.tensor(
|
|
851
|
+
[label_map.get(l, 0) for l in labels],
|
|
852
|
+
dtype=torch.long,
|
|
853
|
+
).unsqueeze(1).to(device)
|
|
854
|
+
except Exception:
|
|
855
|
+
pass
|
|
856
|
+
|
|
857
|
+
with torch.no_grad():
|
|
858
|
+
gen_out = model.module.generative(
|
|
859
|
+
z=z,
|
|
860
|
+
library=library_tensor,
|
|
861
|
+
batch_index=batch_index,
|
|
862
|
+
y=y_tensor,
|
|
863
|
+
)
|
|
864
|
+
px_dist = gen_out["px"]
|
|
865
|
+
vals = (
|
|
866
|
+
px_dist.sample().cpu().numpy()
|
|
867
|
+
if hasattr(px_dist, "sample")
|
|
868
|
+
else px_dist.mean.cpu().numpy()
|
|
869
|
+
)
|
|
870
|
+
|
|
871
|
+
col_names = (
|
|
872
|
+
data.var_names.tolist()
|
|
873
|
+
if hasattr(data, "var_names")
|
|
874
|
+
else [c for c in (data.columns if isinstance(data, pd.DataFrame) else []) if c != target_col]
|
|
875
|
+
)
|
|
876
|
+
synth_df = pd.DataFrame(vals, columns=col_names)
|
|
877
|
+
if target_col is not None and data is not None:
|
|
878
|
+
if hasattr(data, "obs") and target_col in data.obs.columns:
|
|
879
|
+
synth_df[target_col] = data.obs[target_col].values[:n_samples]
|
|
880
|
+
elif isinstance(data, pd.DataFrame) and target_col in data.columns:
|
|
881
|
+
synth_df[target_col] = data[target_col].values[:n_samples]
|
|
882
|
+
return synth_df
|
|
883
|
+
|
|
884
|
+
raise RuntimeError(
|
|
885
|
+
f"decode_from_latent() is not supported for method '{self.method}'. "
|
|
886
|
+
"Supported: tvae, rtvae, scvi, scanvi."
|
|
887
|
+
)
|
|
888
|
+
|
|
645
889
|
@staticmethod
|
|
646
890
|
def to_anndata(
|
|
647
891
|
df: pd.DataFrame,
|
|
@@ -1087,7 +1331,7 @@ class RealGenerator(BaseGenerator):
|
|
|
1087
1331
|
|
|
1088
1332
|
self.synthesizer = syn
|
|
1089
1333
|
self.method = "tvae"
|
|
1090
|
-
self.metadata = {"columns": data.columns.tolist()}
|
|
1334
|
+
self.metadata = {"columns": data.columns.tolist(), "target_col": target_col}
|
|
1091
1335
|
|
|
1092
1336
|
if differentiation_factor > 0.0 and target_col and target_col in data.columns:
|
|
1093
1337
|
synth_df = self.apply_latent_differentiation(
|
|
@@ -1189,17 +1433,22 @@ class RealGenerator(BaseGenerator):
|
|
|
1189
1433
|
result = result.drop(columns=[stage_col])
|
|
1190
1434
|
if custom_distributions:
|
|
1191
1435
|
col = next(iter(custom_distributions), None)
|
|
1192
|
-
if col and col in
|
|
1436
|
+
if col and col in result.columns:
|
|
1193
1437
|
dist = custom_distributions[col]
|
|
1194
1438
|
self.logger.info(
|
|
1195
|
-
f"
|
|
1439
|
+
f" Applying custom_distributions on '{col}' — {dist} (resampling from staged result)"
|
|
1196
1440
|
)
|
|
1197
1441
|
frames = []
|
|
1198
1442
|
for cls, proportion in dist.items():
|
|
1199
1443
|
n_cls = max(1, round(n_samples * proportion))
|
|
1200
|
-
|
|
1201
|
-
|
|
1202
|
-
|
|
1444
|
+
cls_pool = result[result[col] == cls]
|
|
1445
|
+
if len(cls_pool) == 0:
|
|
1446
|
+
cls_pool = result
|
|
1447
|
+
frames.append(
|
|
1448
|
+
cls_pool.sample(n=n_cls, replace=len(cls_pool) < n_cls,
|
|
1449
|
+
random_state=self.random_state)
|
|
1450
|
+
.assign(**{col: cls})
|
|
1451
|
+
)
|
|
1203
1452
|
result = (
|
|
1204
1453
|
pd.concat(frames, ignore_index=True)
|
|
1205
1454
|
.sample(frac=1, random_state=self.random_state)
|
|
@@ -1249,7 +1498,7 @@ class RealGenerator(BaseGenerator):
|
|
|
1249
1498
|
|
|
1250
1499
|
self.synthesizer = syn
|
|
1251
1500
|
self.method = "rtvae"
|
|
1252
|
-
self.metadata = {"columns": data.columns.tolist()}
|
|
1501
|
+
self.metadata = {"columns": data.columns.tolist(), "target_col": target_col}
|
|
1253
1502
|
|
|
1254
1503
|
if differentiation_factor > 0.0 and target_col and target_col in data.columns:
|
|
1255
1504
|
synth_df = self.apply_latent_differentiation(
|
|
@@ -3,10 +3,11 @@ External Reporter Module
|
|
|
3
3
|
Wraps YData Profiling to generate advanced HTML reports.
|
|
4
4
|
"""
|
|
5
5
|
|
|
6
|
+
import logging
|
|
6
7
|
import os
|
|
7
|
-
import pandas as pd
|
|
8
8
|
from typing import Optional
|
|
9
|
-
|
|
9
|
+
|
|
10
|
+
import pandas as pd
|
|
10
11
|
|
|
11
12
|
# Initializing logger
|
|
12
13
|
logger = logging.getLogger("ExternalReporter")
|
|
@@ -83,6 +84,7 @@ class ExternalReporter:
|
|
|
83
84
|
sortby=sortby,
|
|
84
85
|
html={"style": {"full_width": True}},
|
|
85
86
|
progress_bar=False,
|
|
87
|
+
minimal=minimal,
|
|
86
88
|
)
|
|
87
89
|
|
|
88
90
|
profile.to_file(output_path)
|
|
@@ -6,22 +6,24 @@ static report comparing a real dataset with a synthetic one.
|
|
|
6
6
|
Uses YData Profiling for analysis and Plotly for interactive visualizations.
|
|
7
7
|
"""
|
|
8
8
|
|
|
9
|
-
import
|
|
10
|
-
import
|
|
11
|
-
|
|
12
|
-
import warnings
|
|
9
|
+
import contextlib
|
|
10
|
+
import io
|
|
11
|
+
import json
|
|
13
12
|
import logging
|
|
14
|
-
from datetime import datetime
|
|
15
13
|
import os
|
|
16
|
-
import
|
|
17
|
-
import
|
|
18
|
-
import
|
|
14
|
+
import warnings
|
|
15
|
+
from datetime import datetime
|
|
16
|
+
from typing import Any, Dict, List, Optional, Union
|
|
17
|
+
|
|
18
|
+
import numpy as np
|
|
19
|
+
import pandas as pd
|
|
20
|
+
|
|
19
21
|
sc = None
|
|
20
22
|
ad = None
|
|
21
23
|
try:
|
|
22
|
-
from scgft_evaluator import ScGFT_Evaluator
|
|
23
|
-
import scanpy as sc
|
|
24
24
|
import anndata as ad
|
|
25
|
+
import scanpy as sc
|
|
26
|
+
from scgft_evaluator import ScGFT_Evaluator
|
|
25
27
|
SCGFT_AVAILABLE = True
|
|
26
28
|
except ImportError:
|
|
27
29
|
SCGFT_AVAILABLE = False
|
|
@@ -33,12 +35,12 @@ try:
|
|
|
33
35
|
except ImportError:
|
|
34
36
|
SKLEARN_AVAILABLE = False
|
|
35
37
|
|
|
36
|
-
from calm_data_generator.
|
|
37
|
-
from calm_data_generator.reports.
|
|
38
|
-
from calm_data_generator.reports.
|
|
39
|
-
from calm_data_generator.reports.
|
|
40
|
-
from calm_data_generator.reports.
|
|
41
|
-
from calm_data_generator.
|
|
38
|
+
from calm_data_generator.generators.configs import ReportConfig # noqa: E402
|
|
39
|
+
from calm_data_generator.reports.base import BaseReporter # noqa: E402
|
|
40
|
+
from calm_data_generator.reports.DiscriminatorReporter import DiscriminatorReporter # noqa: E402
|
|
41
|
+
from calm_data_generator.reports.ExternalReporter import ExternalReporter # noqa: E402
|
|
42
|
+
from calm_data_generator.reports.LocalIndexGenerator import LocalIndexGenerator # noqa: E402
|
|
43
|
+
from calm_data_generator.reports.Visualizer import Visualizer # noqa: E402
|
|
42
44
|
|
|
43
45
|
# Direct usage of sdmetrics
|
|
44
46
|
try:
|
|
@@ -113,9 +115,9 @@ class QualityReporter(BaseReporter):
|
|
|
113
115
|
return self._assess_quality_scores(real_df, synthetic_df)
|
|
114
116
|
|
|
115
117
|
def calculate_ari(
|
|
116
|
-
self,
|
|
117
|
-
real_df: pd.DataFrame,
|
|
118
|
-
synthetic_df: pd.DataFrame,
|
|
118
|
+
self,
|
|
119
|
+
real_df: pd.DataFrame,
|
|
120
|
+
synthetic_df: pd.DataFrame,
|
|
119
121
|
target_col: str
|
|
120
122
|
) -> Dict[str, float]:
|
|
121
123
|
"""
|
|
@@ -589,10 +591,15 @@ class QualityReporter(BaseReporter):
|
|
|
589
591
|
|
|
590
592
|
# Cross duplicates
|
|
591
593
|
try:
|
|
592
|
-
|
|
593
|
-
|
|
594
|
-
|
|
595
|
-
|
|
594
|
+
shared_cols = [c for c in synthetic_df.columns if c in real_df.columns]
|
|
595
|
+
real_unique = real_df[shared_cols].drop_duplicates()
|
|
596
|
+
synth_cast = synthetic_df[shared_cols].copy()
|
|
597
|
+
for col in shared_cols:
|
|
598
|
+
try:
|
|
599
|
+
synth_cast[col] = synth_cast[col].astype(real_unique[col].dtype)
|
|
600
|
+
except (ValueError, TypeError):
|
|
601
|
+
pass
|
|
602
|
+
merged = synth_cast.merge(real_unique, on=shared_cols, how="left", indicator=True)
|
|
596
603
|
cross_dup_count = (merged["_merge"] == "both").sum()
|
|
597
604
|
except Exception as e:
|
|
598
605
|
self.logger.warning(f"Could not compute cross-duplication count; defaulting to 0. Reason: {e}")
|
|
@@ -755,9 +762,9 @@ class QualityReporter(BaseReporter):
|
|
|
755
762
|
return None
|
|
756
763
|
|
|
757
764
|
def _calculate_ari_metrics(
|
|
758
|
-
self,
|
|
759
|
-
real_df: pd.DataFrame,
|
|
760
|
-
synthetic_df: pd.DataFrame,
|
|
765
|
+
self,
|
|
766
|
+
real_df: pd.DataFrame,
|
|
767
|
+
synthetic_df: pd.DataFrame,
|
|
761
768
|
target_col: Optional[str]
|
|
762
769
|
) -> Optional[Dict[str, float]]:
|
|
763
770
|
"""
|
|
@@ -775,21 +782,21 @@ class QualityReporter(BaseReporter):
|
|
|
775
782
|
features = df.select_dtypes(include=[np.number]).drop(columns=[t_col], errors='ignore')
|
|
776
783
|
if features.empty:
|
|
777
784
|
return None
|
|
778
|
-
|
|
785
|
+
|
|
779
786
|
# Fill NaNs for KMeans
|
|
780
787
|
X = features.fillna(0).values
|
|
781
|
-
|
|
788
|
+
|
|
782
789
|
# Handle cases where we have fewer samples than clusters
|
|
783
790
|
k = min(2, len(X))
|
|
784
791
|
if k < 2:
|
|
785
792
|
return 0.0
|
|
786
|
-
|
|
793
|
+
|
|
787
794
|
kmeans = KMeans(n_clusters=k, n_init=10, random_state=42)
|
|
788
795
|
cluster_labels = kmeans.fit_predict(X)
|
|
789
|
-
|
|
796
|
+
|
|
790
797
|
# Get true labels (convert to categorical codes if necessary)
|
|
791
798
|
true_labels = pd.Categorical(df[t_col]).codes
|
|
792
|
-
|
|
799
|
+
|
|
793
800
|
return float(adjusted_rand_score(true_labels, cluster_labels))
|
|
794
801
|
|
|
795
802
|
ari_real = get_ari(real_df, target_col)
|
|
@@ -805,9 +812,9 @@ class QualityReporter(BaseReporter):
|
|
|
805
812
|
return None
|
|
806
813
|
|
|
807
814
|
def _run_scgft_evaluation(
|
|
808
|
-
self,
|
|
809
|
-
real_df: pd.DataFrame,
|
|
810
|
-
synthetic_df: pd.DataFrame,
|
|
815
|
+
self,
|
|
816
|
+
real_df: pd.DataFrame,
|
|
817
|
+
synthetic_df: pd.DataFrame,
|
|
811
818
|
output_dir: str,
|
|
812
819
|
target_col: Optional[str] = None
|
|
813
820
|
) -> None:
|
|
@@ -830,10 +837,10 @@ class QualityReporter(BaseReporter):
|
|
|
830
837
|
numeric_cols = real_df.select_dtypes(include=[np.number]).columns.tolist()
|
|
831
838
|
if target_col and target_col in numeric_cols:
|
|
832
839
|
numeric_cols.remove(target_col)
|
|
833
|
-
|
|
840
|
+
|
|
834
841
|
adata_real = ad.AnnData(real_df[numeric_cols])
|
|
835
842
|
adata_synth = ad.AnnData(synthetic_df[numeric_cols])
|
|
836
|
-
|
|
843
|
+
|
|
837
844
|
if target_col and target_col in real_df.columns:
|
|
838
845
|
adata_real.obs["cell_type"] = real_df[target_col].astype(str).values
|
|
839
846
|
adata_synth.obs["cell_type"] = synthetic_df[target_col].astype(str).values
|
|
@@ -845,7 +852,7 @@ class QualityReporter(BaseReporter):
|
|
|
845
852
|
# 2. Basic Preprocessing for scvi metrics (PCA is required)
|
|
846
853
|
if self.verbose:
|
|
847
854
|
print(" -> Preprocessing AnnData (PCA)...")
|
|
848
|
-
|
|
855
|
+
|
|
849
856
|
sc.pp.pca(adata_real)
|
|
850
857
|
sc.pp.pca(adata_synth)
|
|
851
858
|
|
|
@@ -878,7 +885,7 @@ class QualityReporter(BaseReporter):
|
|
|
878
885
|
|
|
879
886
|
# 4. Save to HTML Report
|
|
880
887
|
scgft_report_path = os.path.join(output_dir, "scgft_report.html")
|
|
881
|
-
|
|
888
|
+
|
|
882
889
|
html_content = f"""
|
|
883
890
|
<html>
|
|
884
891
|
<head>
|
|
@@ -911,7 +918,7 @@ class QualityReporter(BaseReporter):
|
|
|
911
918
|
</body>
|
|
912
919
|
</html>
|
|
913
920
|
"""
|
|
914
|
-
|
|
921
|
+
|
|
915
922
|
with open(scgft_report_path, "w") as html_file:
|
|
916
923
|
html_file.write(html_content)
|
|
917
924
|
|
{calm_data_generator-2.2.0 → calm_data_generator-2.2.1/calm_data_generator.egg-info}/PKG-INFO
RENAMED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.4
|
|
2
2
|
Name: calm-data-generator
|
|
3
|
-
Version: 2.2.
|
|
3
|
+
Version: 2.2.1
|
|
4
4
|
Summary: CALM-Data-Generator: A Python library for synthetic data generation with support for drift injection and clinical data.
|
|
5
5
|
Author-email: Alejandro Belda Fernandez <alejandrobeldafernandez@gmail.com>
|
|
6
6
|
License: MIT
|
|
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
|
|
|
4
4
|
|
|
5
5
|
[project]
|
|
6
6
|
name = "calm-data-generator"
|
|
7
|
-
version = "2.2.
|
|
7
|
+
version = "2.2.1"
|
|
8
8
|
description = "CALM-Data-Generator: A Python library for synthetic data generation with support for drift injection and clinical data."
|
|
9
9
|
readme = "README.md"
|
|
10
10
|
requires-python = ">=3.10"
|