Initial IMU EKF project

This commit is contained in:
2026-06-16 18:22:26 +08:00
commit 97cd9a56c6
9 changed files with 1946 additions and 0 deletions
+308
View File
@@ -0,0 +1,308 @@
import csv
import math
import tempfile
import unittest
from pathlib import Path
from scripts import run_imu_ekf
class RunImuEkfTests(unittest.TestCase):
def _write_sample_csv(self, path: Path, rows: list[tuple[float, float, float]]):
lines = [
"# odr=0x0F - 500 Hz",
"sensor_uptime_s,temp_c,acc_x_g,acc_y_g,acc_z_g,gyro_x_dps,gyro_y_dps,gyro_z_dps",
]
for sensor_time, gyro_z_dps, acc_z_g in rows:
lines.append(f"{sensor_time},28.0,0,0,{acc_z_g},0,0,{gyro_z_dps}")
path.write_text("\n".join(lines), encoding="utf-8-sig")
def test_csv_parser_skips_metadata_and_reads_required_columns(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "sample.csv"
path.write_text(
"\n".join(
[
"# odr=0x0F - 500 Hz",
"# gyro_bw=0x01 - ODR/4",
"sensor_uptime_s,temp_c,acc_x_g,acc_y_g,acc_z_g,gyro_x_dps,gyro_y_dps,gyro_z_dps,ignored",
"0.000,28.0,0,0,1,1,2,3,x",
"0.002,28.0,0,0,1,1,2,3,x",
]
),
encoding="utf-8-sig",
)
metadata, rows = run_imu_ekf.read_imu_csv(path)
self.assertNotIsInstance(rows, list)
rows = list(rows)
self.assertEqual(metadata["odr"], "0x0F - 500 Hz")
self.assertEqual(metadata["gyro_bw"], "0x01 - ODR/4")
self.assertEqual(len(rows), 2)
self.assertEqual(rows[1].sensor_uptime_s, 0.002)
self.assertEqual(rows[0].gyro_dps[2], 3.0)
def test_csv_parser_rejects_missing_required_columns(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bad.csv"
path.write_text(
"sensor_uptime_s,temp_c,acc_x_g,acc_y_g,acc_z_g,gyro_x_dps,gyro_y_dps\n"
"0.0,28,0,0,1,0,0\n",
encoding="utf-8-sig",
)
with self.assertRaisesRegex(ValueError, "gyro_z_dps"):
run_imu_ekf.read_imu_csv(path)
def test_process_file_writes_expected_result_columns(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(input_path, [(index * 0.002, 0.0, 1.0) for index in range(20)])
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=1)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
self.assertEqual(len(rows), 20)
self.assertIn("roll_deg", rows[0])
self.assertIn("relative_yaw_deg", rows[0])
self.assertIn("segment_id", rows[0])
self.assertIn("gyro_bias_z_dps", rows[0])
self.assertEqual(result.input_rows, 20)
def test_process_file_reads_imu_rows_once(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(input_path, [(0.0, 0.0, 1.0)])
rows = [
run_imu_ekf.ImuRow(0.0, 28.0, (0.0, 0.0, 1.0), (0.0, 0.0, 0.0)),
run_imu_ekf.ImuRow(0.5, 28.0, (0.0, 0.0, 1.0), (0.0, 0.0, 0.0)),
run_imu_ekf.ImuRow(1.0, 28.0, (0.0, 0.0, 1.0), (0.0, 0.0, 0.0)),
]
call_count = 0
original_iter = run_imu_ekf.iter_imu_rows
def single_use_iter(path):
nonlocal call_count
call_count += 1
if call_count > 1:
raise AssertionError("process_file must stream iter_imu_rows once")
return iter(rows)
run_imu_ekf.iter_imu_rows = single_use_iter
try:
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=1)
finally:
run_imu_ekf.iter_imu_rows = original_iter
self.assertEqual(call_count, 1)
self.assertEqual(result.input_rows, 3)
def test_init_seconds_zero_disables_gyro_bias_initialization(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(input_path, [(0.0, 7.5, 1.0), (0.1, 7.5, 1.0)])
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=0)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
self.assertEqual(float(rows[0]["gyro_bias_z_dps"]), 0.0)
self.assertEqual(float(rows[1]["gyro_bias_z_dps"]), 0.0)
def test_init_seconds_rejects_values_outside_integer_range(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
self._write_sample_csv(input_path, [(0.0, 0.0, 1.0)])
with self.assertRaisesRegex(ValueError, "init_seconds.*0.*10.*integer"):
run_imu_ekf.process_file(input_path, Path(tmp) / "out", init_seconds=11)
with self.assertRaisesRegex(ValueError, "init_seconds.*0.*10.*integer"):
run_imu_ekf.process_file(input_path, Path(tmp) / "out", init_seconds=1.5)
def test_relative_yaw_is_unwrapped_in_csv_and_html_samples(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(input_path, [(float(index), 100.0, 1.0) for index in range(5)])
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=0)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
self.assertLess(float(rows[2]["yaw_deg"]), 0.0)
self.assertGreater(float(rows[2]["relative_yaw_deg"]), 180.0)
self.assertGreater(float(rows[4]["relative_yaw_deg"]), 360.0)
self.assertIn("relative_yaw_deg", result.samples[0])
def test_process_file_segments_device_restart_near_zero(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(
input_path,
[(9.9, 0.0, 1.0), (10.0, 0.0, 1.0), (0.002, 0.0, 1.0), (0.004, 0.0, 1.0)],
)
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=0)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
self.assertEqual([row["segment_id"] for row in rows], ["0", "0", "1", "1"])
self.assertEqual(float(rows[2]["dt_s"]), 0.0)
def test_process_file_flushes_short_uninitialized_segment_before_restart(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(
input_path,
[(9.9, 0.0, 1.0), (10.0, 0.0, 1.0), (0.002, 20.0, 1.0), (0.502, 20.0, 1.0), (1.002, 20.0, 1.0)],
)
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=1)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
self.assertEqual(len(rows), 5)
self.assertEqual([row["segment_id"] for row in rows], ["0", "0", "1", "1", "1"])
self.assertEqual(float(rows[2]["dt_s"]), 0.0)
def test_process_file_initializes_each_segment_from_its_own_window(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(
input_path,
[
(0.0, 0.0, 1.0),
(0.5, 0.0, 1.0),
(1.1, 0.0, 1.0),
(0.002, 20.0, 1.0),
(0.502, 20.0, 1.0),
(1.002, 20.0, 1.0),
],
)
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=1)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
segment0_bias = [float(row["gyro_bias_z_dps"]) for row in rows if row["segment_id"] == "0"]
segment1_bias = [float(row["gyro_bias_z_dps"]) for row in rows if row["segment_id"] == "1"]
self.assertTrue(segment0_bias)
self.assertTrue(segment1_bias)
self.assertTrue(all(abs(value) < 0.01 for value in segment0_bias))
self.assertAlmostEqual(segment1_bias[-1], 20.0, delta=0.01)
def test_process_file_uses_restart_initialization_window_not_single_row(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
output_dir = Path(tmp) / "out"
self._write_sample_csv(
input_path,
[
(0.0, 0.0, 1.0),
(0.5, 0.0, 1.0),
(1.1, 0.0, 1.0),
(0.002, 0.0, 1.0),
(0.502, 20.0, 1.0),
(1.002, 20.0, 1.0),
],
)
result = run_imu_ekf.process_file(input_path, output_dir, init_seconds=1)
with result.output_csv.open("r", encoding="utf-8", newline="") as handle:
rows = list(csv.DictReader(handle))
segment1_bias = [float(row["gyro_bias_z_dps"]) for row in rows if row["segment_id"] == "1"]
self.assertAlmostEqual(segment1_bias[-1], 40.0 / 3.0, delta=0.01)
def test_process_file_rejects_timestamp_drop_not_near_zero(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
self._write_sample_csv(input_path, [(10.0, 0.0, 1.0), (9.5, 0.0, 1.0)])
with self.assertRaisesRegex(ValueError, "data row 2.*10.0.*9.5"):
run_imu_ekf.process_file(input_path, Path(tmp) / "out", init_seconds=0)
def test_process_file_rejects_small_nonzero_timestamp_drop(self):
with tempfile.TemporaryDirectory() as tmp:
input_path = Path(tmp) / "imu_sample.csv"
self._write_sample_csv(input_path, [(0.91, 0.0, 1.0), (0.90, 0.0, 1.0)])
with self.assertRaisesRegex(ValueError, "data row 2.*0.91.*0.9"):
run_imu_ekf.process_file(input_path, Path(tmp) / "out", init_seconds=0)
def test_iter_imu_rows_rejects_nan_and_infinity(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "bad.csv"
path.write_text(
"\n".join(
[
"sensor_uptime_s,temp_c,acc_x_g,acc_y_g,acc_z_g,gyro_x_dps,gyro_y_dps,gyro_z_dps",
"0.0,28.0,0,0,nan,0,0,0",
"0.1,28.0,0,0,1,0,0,inf",
]
),
encoding="utf-8-sig",
)
with self.assertRaisesRegex(ValueError, "acc_z_g.*finite.*data row 1"):
list(run_imu_ekf.iter_imu_rows(path))
def test_json_for_script_rejects_nan_and_infinity(self):
with self.assertRaisesRegex(ValueError, "finite"):
run_imu_ekf._json_for_script([{"bad": math.nan}])
with self.assertRaisesRegex(ValueError, "finite"):
run_imu_ekf._json_for_script([{"bad": math.inf}])
def test_html_report_is_self_contained_and_bounded(self):
result = run_imu_ekf.EkfFileResult(
input_csv=Path("imu_sample</script><script>alert(1)</script>.csv"),
output_csv=Path("out/<img src=x onerror=alert(2)>.csv"),
input_rows=3,
metadata={"odr": "<img src=x onerror=alert(3)>", "close": "</script><script>alert(4)</script>"},
samples=[
{
"sensor_uptime_s": 0.0,
"roll_deg": 0.0,
"pitch_deg": 0.0,
"yaw_deg": 0.0,
"relative_yaw_deg": 0.0,
"gyro_bias_x_dps": 0.0,
"gyro_bias_y_dps": 0.0,
"gyro_bias_z_dps": 0.0,
"acc_residual_norm": 0.0,
}
],
)
with tempfile.TemporaryDirectory() as tmp:
html_path = Path(tmp) / "viewer.html"
run_imu_ekf.write_html_report([result], html_path, max_points=1)
html = html_path.read_text(encoding="utf-8")
self.assertIn("<!doctype html>", html.lower())
self.assertIn("application/json", html)
self.assertNotIn("innerHTML", html)
self.assertNotIn("<img src=x", html)
self.assertNotIn("</script><script>alert", html)
self.assertNotIn("window.EKF_DATA", html)
self.assertNotIn("https://", html)
if __name__ == "__main__":
unittest.main()