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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 1 addition & 2 deletions scripts/comparison_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,7 @@
import numpy as np
from conversion.data_formats import MarkImpressionType, MarkStriationType, MarkType
from conversion.surface_comparison.models import ComparisonParams

from scripts.conversion_utils import parse_db_scratch
from conversion_utils import parse_db_scratch

logger = logging.getLogger(__name__)
_MARK_TYPE_FOLDER_MAP: list[tuple[str, MarkType]] = sorted(
Expand Down
11 changes: 5 additions & 6 deletions scripts/convert_scores.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,17 +33,16 @@
from typing import Any

import requests
from conversion.data_formats import MarkImpressionType

from scripts.comparison_utils import (
from comparison_utils import (
ComparisonEntry,
_build_body,
_save_result,
find_all_mark_types,
generate_pairs,
)
from scripts.conversion_utils import ConversionConfig, run_parallel
from scripts.csv_pairs import (
from conversion.data_formats import MarkImpressionType
from conversion_utils import ConversionConfig, run_parallel
from csv_pairs import (
DONE_STATUSES,
CsvTask,
ScoreWriter,
Expand All @@ -53,7 +52,7 @@
find_result_file,
read_pairs_csv,
)
from scripts.http_utils import _cleanup_vault, _post_with_retry, download_urls
from http_utils import _cleanup_vault, _post_with_retry, download_urls

logging.basicConfig(level=logging.WARNING, format="%(levelname)s: %(message)s")
logger = logging.getLogger(__name__)
Expand Down
13 changes: 7 additions & 6 deletions scripts/csv_pairs.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,18 +23,18 @@
from pathlib import Path
from typing import Any

from comparison_utils import ComparisonEntry, infer_mark_type
from conversion.data_formats import MarkImpressionType, MarkType

from scripts.comparison_utils import ComparisonEntry, infer_mark_type
from scripts.conversion_utils import ConversionConfig
from conversion_utils import ConversionConfig

logger = logging.getLogger(__name__)

#: Folder holding the per-mark-type result folders, inside the database folder.
RESULTS_SUBDIR = Path("mark-comparison-results")

#: Keys in the API response. The metrics sit in a nested block
#: (``ComparisonImpressionMetrics`` / ``StriationComparisonResults``).
#: Keys in the API response. Most metrics sit in a nested block
#: (``ComparisonImpressionMetrics`` / ``StriationComparisonResults``); ``n_cells``
#: is a computed field at the top level of ``ComparisonResponseImpression``.
RESULTS_KEY = "comparison_results"
MATCHING_CELLS_KEY = "score"
TOTAL_CELLS_KEY = "n_cells"
Expand Down Expand Up @@ -103,8 +103,9 @@ def extract_metrics(result: dict[str, Any] | None, mark_type: MarkType) -> dict[
return {}

if isinstance(mark_type, MarkImpressionType):
# n_cells is a computed field on the response itself, not part of the nested metrics block.
metrics = {
"total_cells": comparison_results.get(TOTAL_CELLS_KEY),
"total_cells": result.get(TOTAL_CELLS_KEY),
"matching_cells": comparison_results.get(MATCHING_CELLS_KEY),
}
else:
Expand Down
Loading