Skip to content
Open
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
16 changes: 8 additions & 8 deletions src/packages/harp-device/src/harp/device/schema/_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -104,14 +104,14 @@ class ConverterContext:
name: str # yml field key ("__value__" for a whole-register value)
interface_type: Optional[str] # the DSL interfaceType (None = raw/native)
mask: Optional[int] # bit mask, when the value is bit-packed
length: int # register elements this value spans, at least one
length: Optional[int] # declared element count; None when the schema declares none
element: np.dtype # base element dtype of the register, from PayloadType
element_size: int # base element byte size of the register

@property
def span(self) -> int:
"""Byte span of the value, its element count times the element size."""
return self.length * self.element_size
return (self.length or 1) * self.element_size

@property
def member_dtype(self) -> np.dtype:
Expand All @@ -124,8 +124,8 @@ def member_dtype(self) -> np.dtype:

@property
def raw_dtype(self) -> np.dtype:
"""Native passthrough dtype, a sub-array when the value spans more than one element."""
if self.length > 1:
"""Native passthrough dtype, a sub-array when the schema declares an element count."""
if self.length is not None:
return np.dtype((self.element.type, (self.length,)))
return self.element

Expand Down Expand Up @@ -319,7 +319,7 @@ def _resolve_converter(self, ctx: ConverterContext) -> Converter[Any]:
if ctx.mask is not None:
return IdentityConverter(ctx.member_dtype) # bit-field: native slice of the element
if it is None:
if ctx.length > 1:
if ctx.length is not None:
return ArrayConverter(ctx.element, ctx.length) # raw passthrough, sub-array
return IdentityConverter(ctx.element) # raw passthrough, single element
# A known primitive that didn't fit is re-interpreted per field (``{Name}Converter``);
Expand All @@ -343,7 +343,7 @@ def _extension(self, symbol: str, ctx: ConverterContext) -> Converter[Any]:
def _default(self, member: PayloadMember, type_name: str, ctx: ConverterContext) -> Any:
"""The typed default value of the field, or ``_NO_DEFAULT`` when it has none."""
_default_value = member.defaultValue if member.defaultValue is not None else member.minValue
if _default_value is None or ctx.length > 1:
if _default_value is None or ctx.length is not None:
return _NO_DEFAULT
value = float(_default_value.root)
group_mask = self._find_mask(type_name)
Expand Down Expand Up @@ -376,7 +376,7 @@ def _build_field(self, key: str, member: PayloadMember, reg: Register) -> Any:
name=key,
interface_type=it,
mask=member.mask,
length=member.length or 1,
length=member.length,
element=elem,
element_size=elem_size,
)
Expand Down Expand Up @@ -451,7 +451,7 @@ def _new_payload(self, class_name: str, owner: str, reg: Register) -> type:
name="__value__",
interface_type=it,
mask=None,
length=reg.length or 1,
length=reg.length,
element=elem,
element_size=elem_size,
)
Expand Down
24 changes: 24 additions & 0 deletions tests/assets/device.yml
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,30 @@ registers:
PulseDO0:
<<: *pulseDO
address: 43
MultiElementPayload:
address: 44
type: U16
length: 3
access: Write
SingleElementPayload:
address: 45
type: U16
length: 1
access: Write
MixedMemberLength:
address: 46
type: U8
length: 4
access: Write
payloadSpec:
Absent:
offset: 0
Single:
offset: 1
length: 1
Multiple:
offset: 2
length: 2
StartPulse:
address: 100
type: U16
Expand Down
2 changes: 1 addition & 1 deletion tests/data/test_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ def _records(cls, n, seed):
dtype = cls.payload_class.payload_dtype
rng = np.random.default_rng(seed)
raw = rng.integers(0, 128, size=n * dtype.itemsize, dtype=np.uint8)
return raw.view(dtype).copy()
return np.frombuffer(raw, dtype=dtype, count=n).copy()


@pytest.fixture
Expand Down
10 changes: 5 additions & 5 deletions tests/device/converters.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@
from harp.protocol import Converter


class DataConverter(Converter[int]):
"""Maps two raw little-endian signed bytes to and from a Python int.
class DataConverter(Converter[np.int32]):
"""Maps two raw little-endian signed bytes to and from a numpy int32.

Models interfaceType: int over a two-byte sub-region of the CustomMemberConverter payload.
"""
Expand All @@ -18,8 +18,8 @@ def __init__(self) -> None:
self._length = 2
self.dtype = np.dtype((np.uint8, (self._length,)))

def decode_scalar(self, view: np.generic) -> int:
return int.from_bytes(bytes(np.asarray(view).tolist()), "little", signed=True)
def decode_scalar(self, view: np.generic) -> np.int32:
return np.int32(int.from_bytes(bytes(np.asarray(view).tolist()), "little", signed=True))

def decode_batch(self, view: NDArray[np.generic]) -> Any:
return np.array(
Expand All @@ -30,7 +30,7 @@ def decode_batch(self, view: NDArray[np.generic]) -> Any:
dtype=object,
)

def encode_into(self, view: NDArray[np.generic], value: int) -> None:
def encode_into(self, view: NDArray[np.generic], value: np.int32) -> None:
view[...] = np.frombuffer(
int(value).to_bytes(self._length, "little", signed=True), dtype=np.uint8
)
41 changes: 35 additions & 6 deletions tests/device/expected_device.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from numpy.typing import NDArray
from harp.protocol import (
AnonymousPayload,
ArrayConverter,
BitMask,
BoolConverter,
Field,
Expand All @@ -19,6 +20,7 @@
RegisterBase,
RegisterS32,
RegisterU16,
RegisterU16Array,
RegisterU8,
StringConverter,
StructPayload,
Expand All @@ -44,6 +46,7 @@
"CustomMemberConverterPayload",
"BitmaskSplitterPayload",
"PortDIOSetPayload",
"MixedMemberLengthPayload",
"StartPulsePayload",
"StartPulseTrainPayload",
"EncoderModePayload",
Expand All @@ -59,6 +62,9 @@
"PortDIOSet",
"PulseDOPort0",
"PulseDO0",
"MultiElementPayload",
"SingleElementPayload",
"MixedMemberLength",
"StartPulse",
"StartPulseTrain",
"EncoderMode",
Expand Down Expand Up @@ -100,9 +106,7 @@ class AnalogDataPayload(StructPayload[np.float32], length=6):
analog0: np.float32 = Field(IdentityConverter(np.float32))
analog1: np.float32 = Field(IdentityConverter(np.float32), offset=1)
analog2: np.float32 = Field(IdentityConverter(np.float32), offset=2)
accelerometer: NDArray[np.float32] = Field(
IdentityConverter(np.dtype((np.float32, (3,)))), offset=3
)
accelerometer: NDArray[np.float32] = Field(ArrayConverter(np.float32, 3), offset=3)


class ComplexConfigurationPayload(StructPayload[np.uint8], length=17):
Expand All @@ -122,9 +126,7 @@ class VersionPayload(StructPayload[np.uint8], length=32):
firmware_version: HarpVersion = Field(HarpVersionConverter(np.uint8), offset=3)
hardware_version: HarpVersion = Field(HarpVersionConverter(np.uint8), offset=6)
core_id: str = Field(StringConverter(3), offset=9)
interface_hash: NDArray[np.uint8] = Field(
IdentityConverter(np.dtype((np.uint8, (20,)))), offset=12
)
interface_hash: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 20), offset=12)


class CustomPayloadPayload(AnonymousPayload[np.uint32]):
Expand Down Expand Up @@ -159,6 +161,14 @@ class PortDIOSetPayload(AnonymousPayload[np.uint8]):
__value__: PortDigitalIOS = BitMask(enum=PortDigitalIOS)


class MixedMemberLengthPayload(StructPayload[np.uint8], length=4):
"""Represents the payload of the MixedMemberLength register."""

absent: np.uint8 = Field(IdentityConverter(np.uint8))
single: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 1), offset=1)
multiple: NDArray[np.uint8] = Field(ArrayConverter(np.uint8, 2), offset=2)


class StartPulsePayload(StructPayload[np.uint16]):
"""Represents the payload of the StartPulse register."""

Expand Down Expand Up @@ -247,6 +257,22 @@ class PulseDO0(RegisterU16):
address: ClassVar[int] = 43


class MultiElementPayload(RegisterU16Array):
address: ClassVar[int] = 44
length: int = 3


class SingleElementPayload(RegisterU16Array):
address: ClassVar[int] = 45
length: int = 1


class MixedMemberLength(RegisterBase[MixedMemberLengthPayload]):
address: ClassVar[int] = 46
payload_type: ClassVar[PayloadType] = PayloadType.U8
payload_class = MixedMemberLengthPayload


class StartPulse(RegisterBase[StartPulsePayload]):
"""Starts a PWM pulse."""

Expand Down Expand Up @@ -285,6 +311,9 @@ class EncoderMode(RegisterBase[EncoderModeMask]):
41: PortDIOSet,
42: PulseDOPort0,
43: PulseDO0,
44: MultiElementPayload,
45: SingleElementPayload,
46: MixedMemberLength,
100: StartPulse,
101: StartPulseTrain,
103: EncoderMode,
Expand Down
2 changes: 1 addition & 1 deletion tests/device/test_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -574,7 +574,7 @@ def _random_records(dtype, n, seed):
"""
rng = np.random.default_rng(seed)
raw = rng.integers(0, 128, size=n * dtype.itemsize, dtype=np.uint8)
return raw.view(dtype).copy()
return np.frombuffer(raw, dtype=dtype, count=n).copy()


@pytest.mark.parametrize("name", sorted(_device_registers()))
Expand Down