Scikit-learn
human-activity-recognition
wearable
wrist
time-series
cpu
scikit-learn
WISP / src /wisp /ablations.py
Zipeng365's picture
Add files using upload-large-folder tool
80b01cc verified
Raw History Blame Contribute Delete
3.57 kB
# Paper-scoped implementation; see SOURCE_PROVENANCE.json.
from __future__ import annotations
from dataclasses import asdict, dataclass
from typing import Any
from wisp.cis_algorithms import build_cpu_estimator
from wisp.search.grammar import CandidateSpec
@dataclass(frozen=True)
class RecipeSpec:
recipe_id: str
component_tag: str
label: str
use_rocket: bool
use_intervals: bool
use_state: bool
use_hmm: bool
n_kernels: int = 384
n_intervals: int = 24
classifier: str = 'ridge'
use_anchor_stats: bool = False
class_balance: str = 'balanced'
def to_candidate(self, *, seed: int, n_jobs: int) -> CandidateSpec:
operators: list[str] = ['zscore']
if self.use_rocket:
operators.extend(['random_conv', 'ppv_pool'])
if self.use_intervals:
operators.extend(['dyadic_intervals', 'quantile_distribution'])
if self.use_state:
operators.extend(['magnitude', 'jerk', 'spectral', 'autocorr', 'symbolic_transition'])
operators.extend([self.classifier, 'hmm' if self.use_hmm else 'argmax'])
return CandidateSpec(name='state_rocket_interval_hmm', route='cpu', params={'n_kernels': int(self.n_kernels), 'n_intervals': int(self.n_intervals), 'classifier': self.classifier, 'use_hmm': bool(self.use_hmm), 'use_rocket': bool(self.use_rocket), 'use_intervals': bool(self.use_intervals), 'use_state': bool(self.use_state), 'use_anchor_stats': False, 'class_balance': self.class_balance, 'seed': int(seed), 'n_jobs': int(n_jobs)}, operators=operators)
def to_dict(self) -> dict[str, Any]:
return asdict(self)
_COMPONENTS: tuple[tuple[str, bool, bool, bool], ...] = (('R', True, False, False), ('I', False, True, False), ('S', False, False, True), ('RI', True, True, False), ('RS', True, False, True), ('IS', False, True, True), ('RIS', True, True, True))
CIS_ABLATIONS: tuple[RecipeSpec, ...] = tuple((RecipeSpec(recipe_id=f"cis_ablation_{tag}_{('viterbi' if use_hmm else 'direct')}", component_tag=tag, label=f"{tag}+{('H' if use_hmm else 'direct')}", use_rocket=rocket, use_intervals=intervals, use_state=state, use_hmm=use_hmm) for tag, rocket, intervals, state in _COMPONENTS for use_hmm in (False, True)))
CIS_BY_ID = {spec.recipe_id: spec for spec in CIS_ABLATIONS}
FULL_CIS_ID = 'cis_ablation_RIS_viterbi'
def get_recipe_spec(recipe_id: str) -> RecipeSpec:
try:
return CIS_BY_ID[recipe_id]
except KeyError as exc:
raise KeyError(f'Unknown paper recipe {recipe_id!r}; expected one of {sorted(CIS_BY_ID)}') from exc
def build_recipe(recipe_id: str, *, seed: int, n_jobs: int):
spec = get_recipe_spec(recipe_id).to_candidate(seed=seed, n_jobs=n_jobs)
return build_cpu_estimator(spec.to_dict())
def resolved_recipe(recipe_id: str, *, seed: int, n_jobs: int) -> dict[str, Any]:
spec = get_recipe_spec(recipe_id)
return {**spec.to_dict(), 'candidate_spec': spec.to_candidate(seed=seed, n_jobs=n_jobs).to_dict(), 'design_name': 'controlled F6 representation combinations', 'frozen_f6_equivalent': recipe_id == FULL_CIS_ID, 'claim_boundary': 'Seven non-empty R/I/S representation combinations are crossed with direct/H. The design supports predeclared representation-addition, representation-removal and decoder contrasts. Because an empty-representation cell is illegal, it is not an unrestricted complete 2^3 factorial. RIS+H exactly matches the frozen F6 code settings.'}
__all__ = ['RecipeSpec', 'CIS_ABLATIONS', 'CIS_BY_ID', 'FULL_CIS_ID', 'get_recipe_spec', 'build_recipe', 'resolved_recipe']