Skip to content
Draft
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
134 changes: 106 additions & 28 deletions notebooks/convert_complex_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,9 @@ class ColumnSpec:
# Additional dimension values (for wide format):
dimensions: dict[str, str] = field(default_factory=dict) # dimension_name -> value

# Forward fill option
forward_fill: bool = False # If True, empty cells are filled with value above

def is_dimension_column(self) -> bool:
"""Check if this column contains dimension values (long format)."""
return bool(self.dimension_name)
Expand All @@ -42,16 +45,23 @@ class DatasetSchema:
identifier_mappings: dict[str, dict[str, Any]] # identifier -> additional attributes
dimension_names: list[str] # All dimension names used in this dataset
column_specs: dict[int, ColumnSpec] # column_number -> what it contains
skip_rows: int = 1 # Skip explanation/header rows
skip_row_ranges: list[tuple[int, int]] = field(default_factory=list) # [(start, end), ...] inclusive ranges
identifier_column: int = 0 # Column with row identifiers (0-based)
has_header: bool = False # Skip header row by default, use technical col names.
value_column_name: str = 'Value' # Name for the value column in output

def __post_init__(self):
if self.identifier_mappings is None:
self.identifier_mappings = {}
if self.column_specs is None:
self.column_specs = {}
# Convert skip_row_ranges to list of tuples if needed
if self.skip_row_ranges:
self.skip_row_ranges = [tuple(r) if isinstance(r, list) else r
for r in self.skip_row_ranges]

def should_skip_row(self, row_idx: int) -> bool:
"""Check if a row index should be skipped."""
for start, end in self.skip_row_ranges:
if start <= row_idx <= end:
return True
return False
Comment on lines +61 to +64

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ [ruff] <SIM110> reported by reviewdog 🐶
Use return any(start <= row_idx <= end for start, end in self.skip_row_ranges) instead of for loop

Suggested change
for start, end in self.skip_row_ranges:
if start <= row_idx <= end:
return True
return False
return any(start <= row_idx <= end for start, end in self.skip_row_ranges)



@dataclass
Expand All @@ -69,15 +79,40 @@ class DatasetTransformer:
def __init__(self, schema: DatasetSchema):
self.schema = schema

def _apply_forward_fill(self, df: pl.DataFrame) -> pl.DataFrame:
"""Apply forward fill to specified columns."""
forward_fill_cols = [
col_num for col_num, col_spec in self.schema.column_specs.items()
if col_spec.forward_fill
]

if not forward_fill_cols:
return df

# Convert to list of column names
col_names = [f'column_{col_num + 1}' for col_num in forward_fill_cols]

# Apply forward fill to each column
for col_name in col_names:
if col_name in df.columns:
df = df.with_columns(
pl.col(col_name).fill_null(strategy='forward').alias(col_name)
)

return df

def load_and_transform(self, file_path: str) -> pl.DataFrame:
"""Load CSV and transform to normalized format."""

# Load all rows first (no skip_rows parameter)
df = pl.read_csv(
file_path,
skip_rows=self.schema.skip_rows,
has_header=self.schema.has_header,
)

# Apply forward fill before filtering rows
df = self._apply_forward_fill(df)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Bug: Forward fill applies before row filtering causes data leakage

The _apply_forward_fill is applied to all rows before skip_row_ranges filtering occurs. This causes values from rows that should be skipped (e.g., header/explanation rows) to leak into subsequent data rows via forward fill. For example, if row 0 has "Category" text and should be skipped, and row 1 has a null that gets forward-filled, the data row will incorrectly receive the header value "Category". The forward fill operation needs to happen after filtering out skipped rows, not before.

Additional Locations (1)

Fix in Cursor Fix in Web


# Extract identifiers from specified column
identifiers = df.select(
pl.col(f'column_{self.schema.identifier_column + 1}').alias('identifier')
Expand All @@ -87,6 +122,10 @@ def load_and_transform(self, file_path: str) -> pl.DataFrame:
result_rows = []

for row_idx, identifier in enumerate(identifiers):
# Check if this row should be skipped
if self.schema.should_skip_row(row_idx):
continue

if identifier is None or str(identifier).strip() == '':
continue

Expand Down Expand Up @@ -155,7 +194,8 @@ def get_summary(self) -> dict[str, Any]:
'data_columns': len(data_columns),
'dimension_columns': len(dimension_columns),
'dimensions': self.schema.dimension_names,
'identifiers': len(self.schema.identifier_mappings)
'identifiers': len(self.schema.identifier_mappings),
'skip_row_ranges': self.schema.skip_row_ranges
}

if data_columns:
Expand All @@ -177,39 +217,77 @@ def load_config(yaml_file_path: str | Path) -> DatasetConfig:

Returns:
DatasetConfig object with schema and file paths

"""
yaml_file_path = Path(yaml_file_path)

if not yaml_file_path.exists():
raise FileNotFoundError(f"Configuration file not found: {yaml_file_path}")

with open(yaml_file_path, 'r', encoding='utf-8') as f:
with Path.open(yaml_file_path, 'r', encoding='utf-8') as f:
config = yaml.safe_load(f)

# Parse column specs
column_specs = {}
for col_num_str, spec_dict in config.get('column_specs', {}).items():
col_num = int(col_num_str)

# Create ColumnSpec based on whether it's a dimension or data column
if 'dimension_name' in spec_dict and spec_dict['dimension_name']:
# This is a dimension column (long format)
column_specs[col_num] = ColumnSpec(
dimension_name=spec_dict['dimension_name']
)

# First, handle column_spec_templates if they exist
templates = config.get('column_spec_templates', {})

# Then process column_specs
for key, spec_dict in config.get('column_specs', {}).items():
# Handle column ranges (e.g., "4-6")
if isinstance(key, str) and '-' in key:
start, end = map(int, key.split('-'))
col_range = range(start, end + 1)
else:
# This is a data column (wide format)
column_specs[col_num] = ColumnSpec(
quantity=spec_dict.get('quantity', ''),
unit=spec_dict.get('unit', ''),
year=spec_dict.get('year'),
dimensions=spec_dict.get('dimensions', {})
)
col_range = [int(key)]

# Create ColumnSpec for each column in range
for col_num in col_range:
# Check if this spec references a template
if 'template' in spec_dict:
template_name = spec_dict['template']
if template_name not in templates:
raise ValueError(f"Template '{template_name}' not found")
base_spec = templates[template_name].copy()
# Merge with any overrides in spec_dict
base_spec.update({k: v for k, v in spec_dict.items() if k != 'template'})
spec_dict = base_spec # noqa: PLW2901

if spec_dict.get('dimension_name'):
# This is a dimension column (long format)
column_specs[col_num] = ColumnSpec(
dimension_name=spec_dict['dimension_name'],
forward_fill=spec_dict.get('forward_fill', False)
)
else:
# This is a data column (wide format)
column_specs[col_num] = ColumnSpec(
quantity=spec_dict.get('quantity', ''),
unit=spec_dict.get('unit', ''),
year=spec_dict.get('year'),
dimensions=spec_dict.get('dimensions', {}),
forward_fill=spec_dict.get('forward_fill', False)
)

# Parse skip_row_ranges
skip_row_ranges = []
if 'skip_row_ranges' in config:
for range_spec in config['skip_row_ranges']:
if isinstance(range_spec, list) and len(range_spec) == 2:
skip_row_ranges.append((range_spec[0], range_spec[1]))
elif isinstance(range_spec, int):
skip_row_ranges.append((range_spec, range_spec))
# Backward compatibility with old skip_rows
elif 'skip_rows' in config:
skip_rows = config['skip_rows']
if skip_rows > 0:
skip_row_ranges.append((0, skip_rows - 1))

# Create DatasetSchema
schema = DatasetSchema(
row_identifier_name=config['row_identifier_name'],
skip_rows=config.get('skip_rows', 1),
skip_row_ranges=skip_row_ranges,
has_header=config.get('has_header', False),
identifier_column=config.get('identifier_column', 0),
identifier_mappings=config.get('identifier_mappings', {}),
Expand All @@ -236,6 +314,7 @@ def process_dataset(config_path: str | Path, verbose: bool = True) -> pl.DataFra

Returns:
Transformed DataFrame

"""
# Load configuration
if verbose:
Expand Down Expand Up @@ -278,7 +357,6 @@ def process_dataset(config_path: str | Path, verbose: bool = True) -> pl.DataFra


def main():
"""Main entry point for command-line usage."""
parser = argparse.ArgumentParser(
description='Transform peculiar CSV datasets into normalized format using YAML configuration.',
formatter_class=argparse.RawDescriptionHelpFormatter,
Expand Down Expand Up @@ -330,7 +408,7 @@ def main():
else:
print(result_df)

return 0
return 0 # noqa: TRY300

except Exception as e:
print(f"Error: {e}", file=sys.stderr)
Expand All @@ -340,4 +418,4 @@ def main():


if __name__ == "__main__":
sys.exit(main())
sys.exit(main())
Loading