366 lines
16 KiB
Python
366 lines
16 KiB
Python
"""Тесты конвейера, не требующие ROS.
|
||
|
||
Проверяется то, что легко сломать незаметно: разбор CDR, порядок точек,
|
||
выпрямление скоса каналов, геометрия плоскости пути, связность с учётом
|
||
глубины, кодирование грибовидного тела и вставка синтетического предмета.
|
||
|
||
pytest ros2_ws/src/flyguard/test
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
import numpy as np
|
||
import pytest
|
||
|
||
ROOT = Path(__file__).resolve().parents[1]
|
||
if str(ROOT) not in sys.path:
|
||
sys.path.insert(0, str(ROOT))
|
||
|
||
from flyguard.cdr import point_dtype # noqa: E402
|
||
from flyguard.geometry import RailPlane # noqa: E402
|
||
from flyguard.lobula import cluster_by_depth # noqa: E402
|
||
from flyguard.mushroom_body import MushroomBody, MushroomBodyConfig # noqa: E402
|
||
from flyguard.retina import ScanLayout # noqa: E402
|
||
|
||
DATA = ROOT / "data" / "for_hackathon"
|
||
if not DATA.exists():
|
||
DATA = ROOT.parent / "data" / "for_hackathon"
|
||
|
||
|
||
# --------------------------------------------------------------------------- CDR
|
||
|
||
def test_point_dtype_handles_unaligned_timestamp():
|
||
"""У лидара поле timestamp (float64) лежит по смещению 18 — без выравнивания."""
|
||
fields = [("x", 0, 7, 1), ("y", 4, 7, 1), ("z", 8, 7, 1), ("intensity", 12, 7, 1),
|
||
("ring", 16, 4, 1), ("timestamp", 18, 8, 1)]
|
||
dt = point_dtype(fields, 26)
|
||
assert dt.itemsize == 26
|
||
assert dt["timestamp"].itemsize == 8
|
||
|
||
|
||
def test_point_dtype_pads_trailing_gap():
|
||
"""Хвост до point_step добивается паддингом — так устроено само сообщение."""
|
||
assert point_dtype([("x", 0, 7, 1)], 26).itemsize == 26
|
||
|
||
|
||
def test_point_dtype_rejects_overlapping_fields():
|
||
with pytest.raises(ValueError, match="перекрыва"):
|
||
point_dtype([("x", 0, 8, 1), ("y", 4, 7, 1)], 26)
|
||
|
||
|
||
# --------------------------------------------------------------------------- решётка
|
||
|
||
def _layout(n_rings=8, n_az=40, n_echo=2, shift=None):
|
||
el = np.linspace(6.0, -6.0, n_rings)
|
||
shift = np.zeros(n_rings, np.int64) if shift is None else shift
|
||
return ScanLayout(el, -0.1, 2.0, shift, np.zeros(n_rings), n_az, n_echo)
|
||
|
||
|
||
def test_rectification_undoes_channel_skew():
|
||
"""Сдвиг строк должен ровно компенсировать азимутальный сдвиг канала."""
|
||
shift = np.array([-2, -1, 0, 1, 2, 0, -1, 1], np.int64)
|
||
lay = _layout(shift=shift)
|
||
j = np.arange(lay.n_az)
|
||
for h in range(lay.n_rings):
|
||
expected = np.clip(j + shift[h], 0, lay.n_az - 1)
|
||
assert np.array_equal(lay.gather[h], expected)
|
||
assert lay.gather_ok[2].all() # нулевой сдвиг — потерь нет
|
||
|
||
|
||
def test_azimuth_grid_is_monotonic_and_centred():
|
||
lay = _layout(n_az=41)
|
||
d = np.diff(lay.az_grid_deg)
|
||
assert np.allclose(d, -0.1)
|
||
assert lay.dirs.shape == (lay.n_rings, lay.n_az, 3)
|
||
assert np.allclose(np.linalg.norm(lay.dirs, axis=-1), 1.0, atol=1e-5)
|
||
|
||
|
||
def test_forward_direction_is_minus_y():
|
||
"""Азимут 0 смотрит вперёд, а вперёд в системе сенсора — это −Y."""
|
||
lay = ScanLayout(np.array([0.0]), -0.1, 0.0, np.zeros(1, np.int64),
|
||
np.zeros(1), 1, 1)
|
||
assert lay.dirs[0, 0, 1] == pytest.approx(-1.0, abs=1e-6)
|
||
|
||
|
||
# --------------------------------------------------------------------------- геометрия
|
||
|
||
def test_floor_range_matches_flat_plane():
|
||
"""Луч вниз под углом θ проходит до полотна H/sin θ — это дальность вдоль
|
||
луча, а не горизонтальное расстояние: именно её сравнивают с измеренной."""
|
||
height = 2.0
|
||
plane = RailPlane(a=0.0, b=0.0, c=-height, inliers=0, rms=0.0)
|
||
for el in (-1.0, -3.0, -10.0):
|
||
lay = ScanLayout(np.array([el]), -0.1, 0.0, np.zeros(1, np.int64),
|
||
np.zeros(1), 1, 1)
|
||
got = plane.floor_range(lay)[0, 0]
|
||
assert got == pytest.approx(height / np.sin(np.radians(-el)), rel=1e-3)
|
||
|
||
|
||
def test_upward_rays_never_hit_floor():
|
||
plane = RailPlane(a=0.0, b=0.0, c=-2.0, inliers=0, rms=0.0)
|
||
lay = ScanLayout(np.array([5.0]), -0.1, 0.0, np.zeros(1, np.int64), np.zeros(1), 1, 1)
|
||
assert not np.isfinite(plane.floor_range(lay)[0, 0])
|
||
|
||
|
||
def test_height_above_plane_accounts_for_tilt():
|
||
plane = RailPlane(a=0.01, b=-0.02, c=-1.5, inliers=0, rms=0.0)
|
||
d = np.array([10.0]); u = np.array([2.0]); z = np.array([0.0])
|
||
expected = 0.0 - (0.01 * 10.0 + (-0.02) * 2.0 + (-1.5))
|
||
assert plane.height_of(d, u, z)[0] == pytest.approx(expected)
|
||
|
||
|
||
# --------------------------------------------------------------------------- кластеризация
|
||
|
||
def test_depth_aware_clustering_splits_on_range_gap():
|
||
"""Предмет перед далёкой стеной не должен слипнуться с ней в одно пятно."""
|
||
mask = np.zeros((4, 10), bool)
|
||
mask[:, :] = True
|
||
r = np.full((4, 10), 150.0, np.float32)
|
||
r[:, 3:6] = 55.0 # предмет на 55 м на фоне 150 м
|
||
labels, n = cluster_by_depth(mask, r, col_reach=1, row_reach=1)
|
||
assert n >= 2
|
||
assert len(set(labels[:, 3:6].ravel())) == 1
|
||
assert labels[0, 0] != labels[0, 4]
|
||
|
||
|
||
def test_clustering_tolerance_grows_with_range():
|
||
"""Допуск относительный: разрыв 1 м слитен на 150 м и разделим на 5 м."""
|
||
mask = np.ones((1, 6), bool)
|
||
far = np.array([[150.0, 150.0, 151.0, 151.0, 150.0, 150.0]], np.float32)
|
||
near = np.array([[5.0, 5.0, 6.0, 6.0, 5.0, 5.0]], np.float32)
|
||
_, n_far = cluster_by_depth(mask, far, col_reach=1, row_reach=0)
|
||
_, n_near = cluster_by_depth(mask, near, col_reach=1, row_reach=0)
|
||
assert n_far == 1
|
||
assert n_near >= 2
|
||
|
||
|
||
def test_empty_mask_is_handled():
|
||
labels, n = cluster_by_depth(np.zeros((3, 3), bool), np.zeros((3, 3), np.float32))
|
||
assert n == 0 and labels.sum() == 0
|
||
|
||
|
||
# --------------------------------------------------------------------------- память
|
||
|
||
def test_mushroom_body_code_is_sparse_and_deterministic():
|
||
mb = MushroomBody(MushroomBodyConfig(n_kc=2000, sparsity=0.01, seed=1))
|
||
X = np.random.default_rng(0).normal(size=(5, mb.n_pn)).astype(np.float32)
|
||
mb.fit_normalizer(X)
|
||
a, b = mb.encode(X), mb.encode(X)
|
||
assert a.shape == (5, mb.n_active)
|
||
assert mb.n_active == 20
|
||
assert np.array_equal(a, b)
|
||
assert len(set(a[0].tolist())) == mb.n_active # без повторов
|
||
|
||
|
||
def test_learning_suppresses_only_what_was_shown():
|
||
mb = MushroomBody(MushroomBodyConfig(n_kc=4000, sparsity=0.01, seed=2))
|
||
rng = np.random.default_rng(3)
|
||
familiar = rng.normal(size=(200, mb.n_pn)).astype(np.float32)
|
||
novel = (rng.normal(size=(50, mb.n_pn)) + 8.0).astype(np.float32)
|
||
mb.fit_normalizer(familiar)
|
||
mb.learn(familiar, rate=0.3)
|
||
assert mb.novelty(familiar).mean() < mb.novelty(novel).mean()
|
||
|
||
|
||
def test_auto_rate_shrinks_with_sample_size():
|
||
mb = MushroomBody(MushroomBodyConfig(n_kc=10_000, sparsity=0.001))
|
||
assert mb.auto_rate(1_000) > mb.auto_rate(100_000)
|
||
assert 0 < mb.auto_rate(10 ** 7) <= 0.5
|
||
|
||
|
||
def test_feature_count_mismatch_is_explicit():
|
||
mb = MushroomBody()
|
||
with pytest.raises(ValueError, match="признак"):
|
||
mb.encode(np.zeros((1, mb.n_pn + 1), np.float32))
|
||
|
||
|
||
# --------------------------------------------------------- фигура и фон, привыкание
|
||
|
||
def _cand(**kw):
|
||
from flyguard.lobula import Candidate
|
||
|
||
base = dict(d=60.0, u=0.1, h=1.0, d_min=59.7, h_min=0.3, width=0.4, height=1.0,
|
||
depth=0.5, containment=0.9, n_rays=30, n_rings=6, n_cols=5,
|
||
gap=8.0, on=0.001, floor_deficit=0.2, shadow=0.2, inten=40.0,
|
||
az_deg=0.5, el_deg=-1.0, bbox=(10, 16, 20, 25))
|
||
base.update(kw)
|
||
return Candidate(**base)
|
||
|
||
|
||
def test_contrast_split_keeps_step_and_drops_smooth_wall():
|
||
"""Гладкая стена не даёт фигуры, ступенька на ней — даёт."""
|
||
from flyguard.lobula import cluster_by_depth, split_by_figure
|
||
|
||
# стена: дальность плавно растёт вдоль строки от 20 до 120 м
|
||
r = np.tile(np.linspace(20.0, 120.0, 200, dtype=np.float32), (12, 1))
|
||
mask = np.ones(r.shape, bool)
|
||
labels, n = cluster_by_depth(mask, r, col_reach=3)
|
||
assert n == 1 # стена связна целиком
|
||
|
||
# контраст гладкой стены равен нулю — резать нечего
|
||
flat = np.zeros_like(r)
|
||
out, n_out = split_by_figure(labels, n, r, flat, max_depth=15.0, thr=6.0)
|
||
assert n_out == n and out is labels
|
||
|
||
# предмет: ступенька, торчащая из стены на 10 м
|
||
gap = np.zeros_like(r)
|
||
gap[4:8, 90:98] = 10.0
|
||
out, n_out = split_by_figure(labels, n, r, gap, max_depth=15.0, thr=6.0)
|
||
kept = out[out > 0]
|
||
assert kept.size == 32 # ровно лучи ступеньки
|
||
assert np.unique(kept).size == 1 # и это одна компонента
|
||
|
||
|
||
def test_habituation_suppresses_shape_repeating_at_different_places():
|
||
from flyguard.mushroom_body import Habituation, HabituationConfig
|
||
|
||
hab = Habituation(HabituationConfig(warmup=0, norm_n=1))
|
||
seen = []
|
||
s = 0.0
|
||
for _ in range(14):
|
||
s += 25.0
|
||
c = _cand(d=60.0)
|
||
hab.update([c], s, 25.0)
|
||
seen.append(c.novelty)
|
||
assert seen[0] > 0.9 and seen[-1] < 0.25
|
||
assert hab.places >= 10
|
||
|
||
|
||
def test_habituation_keeps_object_standing_in_one_place_novel():
|
||
"""Подъезд к неподвижному предмету — это одно место, а не сто повторов."""
|
||
from flyguard.mushroom_body import Habituation, HabituationConfig
|
||
|
||
hab = Habituation(HabituationConfig(warmup=0, norm_n=1))
|
||
s = 0.0
|
||
nov = []
|
||
for k in range(50):
|
||
s += 3.0
|
||
nov.append(_cand(d=160.0 - 3.0 * k))
|
||
hab.update([nov[-1]], s, 3.0)
|
||
assert hab.places <= 2 # одно место (плюс дрейф оценки)
|
||
assert nov[-1].novelty > 0.7
|
||
|
||
|
||
def test_habituation_forgets_after_a_long_run():
|
||
from flyguard.mushroom_body import Habituation, HabituationConfig
|
||
|
||
hab = Habituation(HabituationConfig(warmup=0, norm_n=1, recover_m=100.0))
|
||
s = 0.0
|
||
for _ in range(14):
|
||
s += 25.0
|
||
hab.update([_cand(d=60.0)], s, 25.0)
|
||
low = hab.level
|
||
hab.advance(600.0)
|
||
assert hab.level < low * 0.2
|
||
|
||
|
||
# --------------------------------------------------------------- считывание MBON
|
||
|
||
def test_mbon_readout_learns_and_round_trips(tmp_path):
|
||
"""Обучается, сохраняется без потерь и даёт калиброванную вероятность."""
|
||
from flyguard.mbon_readout import MbonConfig, MbonReadout
|
||
|
||
rng = np.random.default_rng(3)
|
||
X = rng.standard_normal((2000, 23)).astype(np.float32)
|
||
y = ((X[:, 0] + 0.5 * X[:, 3]) > 0.3).astype(np.float32)
|
||
m = MbonReadout(MbonConfig(n_kc=4000, sparsity=0.02), n_pn=23)
|
||
m.fit_normalizer(X)
|
||
m.learn(X, y, epochs=60, lr=4.0)
|
||
|
||
s = m.score(X)
|
||
assert s.min() >= 0.0 and s.max() <= 1.0
|
||
assert ((s > 0.5) == (y > 0.5)).mean() > 0.8 # калибровка, не только порядок
|
||
|
||
path = tmp_path / "mbon.npz"
|
||
m.save(path)
|
||
again = MbonReadout.load(path)
|
||
assert np.allclose(again.score(X[:64]), s[:64], atol=1e-6)
|
||
|
||
|
||
def test_mbon_replaces_hand_formula_and_blend_interpolates():
|
||
"""Модель входит в вес улики, а смешивание даёт обе крайности."""
|
||
from flyguard.central_complex import _quality
|
||
|
||
c = _cand(gap=0.0, containment=0.3, depth=9.0, h_min=1.4, n_rays=5)
|
||
hand = _quality(c) # без модели
|
||
c.extra["mbon"] = 0.99
|
||
assert _quality(c, mbon_blend=1.0) > hand * 3 # модель вытягивает
|
||
assert _quality(c, mbon_blend=0.0) == hand # ручная формула нетронута
|
||
mid = _quality(c, mbon_blend=0.5)
|
||
assert hand < mid < _quality(c, mbon_blend=1.0)
|
||
|
||
# уверенное «это обстановка» гасит даже хорошую геометрию
|
||
good = _cand(gap=12.0, containment=1.0, depth=0.4, h_min=0.3, n_rays=80)
|
||
strong = _quality(good)
|
||
good.extra["mbon"] = 0.01
|
||
assert _quality(good, mbon_blend=1.0) < strong * 0.2
|
||
|
||
|
||
def test_mbon_absent_leaves_pipeline_unchanged():
|
||
"""Без модели вес улики считается ровно как раньше."""
|
||
from flyguard.central_complex import _quality
|
||
|
||
c = _cand()
|
||
assert "mbon" not in c.extra
|
||
assert _quality(c, mbon_blend=1.0) == _quality(c, mbon_blend=0.0)
|
||
|
||
|
||
# ------------------------------------------------------- модель сенсора: яркость
|
||
|
||
def test_injected_intensity_copies_the_surroundings():
|
||
"""Яркость вставки берётся из записи, а не назначается.
|
||
|
||
Абсолютной шкалы интенсивности в данных нет: медиана кандидатов по бэгам
|
||
3…7, а в записи с настоящим предметом 23.5. Любое назначенное число делает
|
||
вставку опознаваемой по одной яркости — сначала как «тускло = предмет»
|
||
(ламбертова ρ·cosθ/r²), потом как «ярко = предмет» (постоянные 35).
|
||
Поэтому у функции нет аргументов ни дальности, ни отражения.
|
||
"""
|
||
import inspect
|
||
|
||
from flyguard.synth import local_intensity
|
||
|
||
# отражения среди аргументов нет; дальность есть, но только как выбор
|
||
# полосы реальных возвратов для лучей, ушедших в пустоту
|
||
args = set(inspect.signature(local_intensity).parameters)
|
||
assert not args & {"reflectivity", "rho", "cos_inc"}
|
||
|
||
rng = np.random.default_rng(0)
|
||
prev = np.array([10.0, 20.0, 0.0, 0.0], np.float32) # два луча в пустоту
|
||
pool = np.full(50, 30.0, np.float32)
|
||
frame = np.full(10_000, 7.0, np.float32)
|
||
v = local_intensity(prev, pool, frame, rng)
|
||
|
||
assert 8.0 < v[0] < 13.0 and 16.0 < v[1] < 25.0 # взято с тех же лучей
|
||
assert 24.0 < v[2] < 38.0 and 24.0 < v[3] < 38.0 # взято из запаса
|
||
|
||
# запас пуст — остаётся кадр целиком
|
||
v2 = local_intensity(np.zeros(200, np.float32), np.zeros(0, np.float32),
|
||
frame, rng)
|
||
assert 5.5 < float(np.median(v2)) < 9.0
|
||
|
||
# без разложения окружения по дальности сама дальность ничего не меняет
|
||
a = local_intensity(prev, pool, frame, np.random.default_rng(3), d=20.0)
|
||
b = local_intensity(prev, pool, frame, np.random.default_rng(3), d=200.0)
|
||
assert np.allclose(a, b)
|
||
|
||
|
||
# --------------------------------------------------------------------------- данные
|
||
|
||
@pytest.mark.skipif(not DATA.exists(), reason="датасет не распакован")
|
||
def test_real_bag_projects_without_angular_error():
|
||
"""На реальном бэге выпрямленная решётка обязана описывать лучи точно."""
|
||
from flyguard.bag import Bag
|
||
|
||
bag = Bag(next(p for p in DATA.iterdir() if p.is_dir()))
|
||
clouds = [pc for _, pc in bag.frames(start=2, stop=8)]
|
||
lay = ScanLayout.calibrate(clouds)
|
||
img = lay.project(clouds[-1])
|
||
|
||
assert img.shape == (lay.n_rings, lay.n_az)
|
||
assert 0.2 < img.valid.mean() < 0.9
|
||
r = img.r_near[img.valid]
|
||
assert r.min() > 0 and r.max() < 250
|
||
assert np.all(img.r_far[img.valid] >= img.r_near[img.valid] - 1e-3)
|