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.
Files changed (147) hide show
  1. {calm_data_generator-2.2.0/calm_data_generator.egg-info → calm_data_generator-2.2.1}/PKG-INFO +1 -1
  2. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/__init__.py +5 -5
  3. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/RealBlockGenerator.py +8 -7
  4. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/RealGenerator.py +256 -7
  5. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/ExternalReporter.py +4 -2
  6. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/QualityReporter.py +46 -39
  7. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1/calm_data_generator.egg-info}/PKG-INFO +1 -1
  8. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/pyproject.toml +1 -1
  9. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_accessors.py +98 -21
  10. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/LICENSE +0 -0
  11. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/MANIFEST.in +0 -0
  12. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/README.md +0 -0
  13. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/cli.py +0 -0
  14. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/API.md +0 -0
  15. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/API_ES.md +0 -0
  16. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CAUSAL_ENGINE_REFERENCE.md +0 -0
  17. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CAUSAL_ENGINE_REFERENCE_ES.md +0 -0
  18. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_BLOCK_GENERATOR_REFERENCE.md +0 -0
  19. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
  20. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_GENERATOR_REFERENCE.md +0 -0
  21. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/CLINICAL_GENERATOR_REFERENCE_ES.md +0 -0
  22. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/COMPLEX_GENERATOR_REFERENCE.md +0 -0
  23. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/COMPLEX_GENERATOR_REFERENCE_ES.md +0 -0
  24. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DOCUMENTATION.md +0 -0
  25. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DOCUMENTATION_ES.md +0 -0
  26. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DRIFT_INJECTOR_REFERENCE.md +0 -0
  27. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/DRIFT_INJECTOR_REFERENCE_ES.md +0 -0
  28. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/PRESETS_REFERENCE.md +0 -0
  29. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/PRESETS_REFERENCE_ES.md +0 -0
  30. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_BLOCK_GENERATOR_REFERENCE.md +0 -0
  31. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
  32. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_GENERATOR_REFERENCE.md +0 -0
  33. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REAL_GENERATOR_REFERENCE_ES.md +0 -0
  34. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REPORTS_REFERENCE.md +0 -0
  35. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/REPORTS_REFERENCE_ES.md +0 -0
  36. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/SCENARIO_INJECTOR_REFERENCE.md +0 -0
  37. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/SCENARIO_INJECTOR_REFERENCE_ES.md +0 -0
  38. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_BLOCK_GENERATOR_REFERENCE.md +0 -0
  39. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_BLOCK_GENERATOR_REFERENCE_ES.md +0 -0
  40. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_GENERATOR_REFERENCE.md +0 -0
  41. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/docs/STREAM_GENERATOR_REFERENCE_ES.md +0 -0
  42. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/__init__.py +0 -0
  43. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/base.py +0 -0
  44. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/Clinic.py +0 -0
  45. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/ClinicGeneratorBlock.py +0 -0
  46. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/ClinicReporter.py +0 -0
  47. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/clinical/__init__.py +0 -0
  48. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/complex/ComplexGenerator.py +0 -0
  49. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/complex/__init__.py +0 -0
  50. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/configs.py +0 -0
  51. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/drift/DriftInjector.py +0 -0
  52. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/drift/__init__.py +0 -0
  53. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/CausalEngine.py +0 -0
  54. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/ScenarioInjector.py +0 -0
  55. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/dynamics/__init__.py +0 -0
  56. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/persistence_models.py +0 -0
  57. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/GeneratorFactory.py +0 -0
  58. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamBlockGenerator.py +0 -0
  59. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamGenerator.py +0 -0
  60. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/StreamReporter.py +0 -0
  61. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/stream/__init__.py +0 -0
  62. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/CustomPluginAdapter.py +0 -0
  63. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/QualityReporter.py +0 -0
  64. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/tabular/__init__.py +0 -0
  65. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/utils/__init__.py +0 -0
  66. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/generators/utils/propagation.py +0 -0
  67. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/logger.py +0 -0
  68. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/BalancePreset.py +0 -0
  69. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ConceptDriftPreset.py +0 -0
  70. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/CopulaPreset.py +0 -0
  71. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DataQualityAuditPreset.py +0 -0
  72. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DiffusionPreset.py +0 -0
  73. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/DriftScenarioPreset.py +0 -0
  74. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/FastPreset.py +0 -0
  75. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/FastPrototypePreset.py +0 -0
  76. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/GradualDriftPreset.py +0 -0
  77. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/HighFidelityPreset.py +0 -0
  78. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ImbalancePreset.py +0 -0
  79. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/LongitudinalHealthPreset.py +0 -0
  80. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/OmicsIntegrationPreset.py +0 -0
  81. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/RareDiseasePreset.py +0 -0
  82. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/ScenarioInjectorPreset.py +0 -0
  83. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/SeasonalTimeSeriesPreset.py +0 -0
  84. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/SingleCellQualityPreset.py +0 -0
  85. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/TimeSeriesPreset.py +0 -0
  86. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/__init__.py +0 -0
  87. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/presets/base.py +0 -0
  88. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/DiscriminatorReporter.py +0 -0
  89. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/LocalIndexGenerator.py +0 -0
  90. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/Visualizer.py +0 -0
  91. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/__init__.py +0 -0
  92. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/reports/base.py +0 -0
  93. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/advanced_drifts.py +0 -0
  94. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinic_generator.py +0 -0
  95. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinical_block_generator.py +0 -0
  96. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/clinical_generator.py +0 -0
  97. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/correlation_drift.py +0 -0
  98. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/drift_injector.py +0 -0
  99. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/real_block_generator.py +0 -0
  100. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/real_generator.py +0 -0
  101. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/reports_deep_dive.py +0 -0
  102. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/scenario_injector.py +0 -0
  103. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/stream_block_generator.py +0 -0
  104. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/stream_generator.py +0 -0
  105. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator/tutorials/tutorial_advanced_methods.py +0 -0
  106. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/SOURCES.txt +0 -0
  107. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/dependency_links.txt +0 -0
  108. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/entry_points.txt +0 -0
  109. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/requires.txt +0 -0
  110. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/calm_data_generator.egg-info/top_level.txt +0 -0
  111. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/requirements.txt +0 -0
  112. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/setup.cfg +0 -0
  113. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_anndata_support.py +0 -0
  114. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_block_generators_extended.py +0 -0
  115. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_causal_engine.py +0 -0
  116. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinic_block_generator.py +0 -0
  117. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinical_advanced.py +0 -0
  118. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_clinical_regression.py +0 -0
  119. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_complex_generator.py +0 -0
  120. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_comprehensive.py +0 -0
  121. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_comprehensive_reporting.py +0 -0
  122. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_correlation_propagation.py +0 -0
  123. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_custom_plugin_adapter.py +0 -0
  124. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_differentiation_factor.py +0 -0
  125. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_discriminator_reporter.py +0 -0
  126. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_disease_effects_fix.py +0 -0
  127. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_correlations.py +0 -0
  128. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_injector_advanced.py +0 -0
  129. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_drift_injector_math.py +0 -0
  130. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_functional_drift.py +0 -0
  131. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_imbalance.py +0 -0
  132. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_migration.py +0 -0
  133. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_presets.py +0 -0
  134. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_quality_metrics.py +0 -0
  135. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_real_block_generator.py +0 -0
  136. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_real_generator_persistence.py +0 -0
  137. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_reporters_extended.py +0 -0
  138. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_reporting_fix.py +0 -0
  139. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_river_integration.py +0 -0
  140. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scenario_extended.py +0 -0
  141. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scgft_reporter.py +0 -0
  142. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_scvi_quality_regression.py +0 -0
  143. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_single_call.py +0 -0
  144. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_stream_block_generator.py +0 -0
  145. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_tabular_extended.py +0 -0
  146. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_timeseries_extended.py +0 -0
  147. {calm_data_generator-2.2.0 → calm_data_generator-2.2.1}/tests/test_timeseries_real.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: calm-data-generator
3
- Version: 2.2.0
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.generators.tabular import RealGenerator, QualityReporter
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 presets
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.0"
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 Optional, Dict, Any, List, Union
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 .RealGenerator import RealGenerator
31
- from calm_data_generator.reports.QualityReporter import QualityReporter
32
- from calm_data_generator.generators.configs import DriftConfig, ReportConfig
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 data.columns:
1436
+ if col and col in result.columns:
1193
1437
  dist = custom_distributions[col]
1194
1438
  self.logger.info(
1195
- f" Generating conditionally per class on '{col}' — {dist}"
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
- cls_df = syn.generate(count=n_cls, random_state=self.random_state).dataframe()
1201
- cls_df[col] = cls
1202
- frames.append(cls_df)
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
- import logging
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 pandas as pd
10
- import numpy as np
11
- from typing import Optional, Dict, Any, List, Union
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 json
17
- import io
18
- import contextlib
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.reports.ExternalReporter import ExternalReporter
37
- from calm_data_generator.reports.Visualizer import Visualizer
38
- from calm_data_generator.reports.LocalIndexGenerator import LocalIndexGenerator
39
- from calm_data_generator.reports.base import BaseReporter
40
- from calm_data_generator.reports.DiscriminatorReporter import DiscriminatorReporter
41
- from calm_data_generator.generators.configs import ReportConfig
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
- real_unique = real_df.drop_duplicates()
593
- merged = synthetic_df.merge(
594
- real_unique, on=list(synthetic_df.columns), how="left", indicator=True
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
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: calm-data-generator
3
- Version: 2.2.0
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.0"
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"