CIRCT 24.0.0git
Loading...
Searching...
No Matches
test_esi.py
Go to the documentation of this file.
1from __future__ import annotations
2
3import os
4from pathlib import Path
5import sys
6import time
7from typing import Optional
8
9import esiaccel as esi
10from esiaccel.accelerator import AcceleratorConnection
11from esiaccel.cosim.pytest import cosim_test
12from esiaccel.types import MMIORegion, MetricPort
13
14HW_DIR = Path(__file__).resolve().parent.parent / "hw"
15
16
17def run(conn: AcceleratorConnection, platform: str = "cosim") -> None:
18 mmio = conn.get_service_mmio()
19 data = mmio.read(8)
20 assert data == 0x207D98E5E5100E51
21
22 assert conn.sysinfo().esi_version() == 0
23 m = conn.manifest()
24 assert m.api_version == 0
25 print(m.type_table)
26
27 # Test the cycle count and clock frequency APIs
28 sysinfo = conn.sysinfo()
29 cycle_count = sysinfo.cycle_count()
30 if cycle_count is not None:
31 print(f"Cycle count: {cycle_count}")
32 assert cycle_count > 0, f"Cycle count should be positive, got {cycle_count}"
33
34 # Test that cycle count is monotonically increasing
35 time.sleep(0.01) # Small delay to let simulation advance
36 cycle_count2 = sysinfo.cycle_count()
37 print(f"Cycle count after delay: {cycle_count2}")
38 assert cycle_count2 > cycle_count, \
39 f"Cycle count should be monotonically increasing: {cycle_count2} <= {cycle_count}"
40
41 # Test again to ensure consistency
42 time.sleep(0.01)
43 cycle_count3 = sysinfo.cycle_count()
44 print(f"Cycle count after second delay: {cycle_count3}")
45 assert cycle_count3 > cycle_count2, \
46 f"Cycle count should be monotonically increasing: {cycle_count3} <= {cycle_count2}"
47 else:
48 print("Cycle count: not available")
49
50 clock_freq = sysinfo.core_clock_frequency()
51 print(f"Clock frequency: {clock_freq} Hz")
52 if platform == "cosim":
53 assert clock_freq == 20_000_000, \
54 f"Expected clock frequency 20_000_000 Hz for cosim, got {clock_freq}"
55 else:
56 if clock_freq is not None:
57 print(f"Core clock frequency: {clock_freq} Hz")
58 assert clock_freq > 0, \
59 f"Clock frequency should be positive, got {clock_freq}"
60 else:
61 print("Core clock frequency: not available")
62
63 d = conn.build_accelerator()
64
65 mmio_svc: esi.accelerator.MMIO
66 for svc in d.services:
67 if isinstance(svc, esi.accelerator.MMIO):
68 mmio_svc = svc
69 break
70
71 for id, region in mmio_svc.regions.items():
72 print(f"Region {id}: {region.base} - {region.base + region.size}")
73
74 def count_telemetry_counters(module) -> int:
75 local_count = sum(
76 isinstance(port, MetricPort) for port in module.ports.values())
77 return local_count + sum(
78 count_telemetry_counters(child) for child in module.children.values())
79
80 num_telemetry_counters = count_telemetry_counters(d)
81 assert num_telemetry_counters > 0
82 # Note: the keys of 'regions' are AppIDPaths, whose repr is the dotted path
83 # without decoration (unlike AppID, which is wrapped in angle brackets).
84 telemetry_region = next((region for id, region in mmio_svc.regions.items()
85 if str(id) == "__telemetry_mmio"), None)
86 assert telemetry_region is not None, \
87 f"no __telemetry_mmio region in {[str(i) for i in mmio_svc.regions]}"
88 telemetry_bytes = num_telemetry_counters * 8
89 expected_telemetry_allocation = 1 << (telemetry_bytes - 1).bit_length()
90 assert telemetry_region.size == expected_telemetry_allocation
91 assert len(mmio_svc.regions) == 6
92
93 ##############################################################################
94 # MMIOClient tests
95 ##############################################################################
96
97 def read_offset(mmio_x: MMIORegion, offset: int, add_amt: int):
98 data = mmio_x.read(offset)
99 if data == add_amt + offset:
100 print(f"PASS: read_offset({offset}, {add_amt}) -> {data}")
101 else:
102 assert False, f"read_offset({offset}, {add_amt}) -> {data}"
103
104 mmio4 = d.ports[esi.AppID("mmio_client", 4)]
105 assert mmio4.descriptor.size == 0x100
106 assert mmio4.descriptor.base % mmio4.descriptor.size == 0
107 read_offset(mmio4, 0, 4)
108 read_offset(mmio4, 13, 4)
109
110 mmio9 = d.ports[esi.AppID("mmio_client", 9)]
111 assert mmio9.descriptor.size == 0x2000
112 assert mmio9.descriptor.base % mmio9.descriptor.size == 0
113 read_offset(mmio9, 0, 9)
114 read_offset(mmio9, 13, 9)
115 read_offset(mmio9, 0x1000, 9)
116 read_offset(mmio9, 0x1200, 9)
117
118 mmio14 = d.ports[esi.AppID("mmio_client", 14)]
119 assert mmio14.descriptor.size == 0x100
120 assert mmio14.descriptor.base % mmio14.descriptor.size == 0
121 read_offset(mmio14, 0, 14)
122 read_offset(mmio14, 13, 14)
123
124 ##############################################################################
125 # MMIOReadWriteClient tests
126 ##############################################################################
127
128 mmio_rw = d.ports[esi.AppID("mmio_rw_client")]
129 assert mmio_rw.descriptor.size == 0x200
130 assert mmio_rw.descriptor.base % mmio_rw.descriptor.size == 0
131
132 def read_offset_check(i: int, add_amt: int):
133 d = mmio_rw.read(i)
134 if d == i + add_amt:
135 print(f"PASS: read_offset_check({i}): {d}")
136 else:
137 assert False, f": read_offset_check({i}): {d}"
138
139 add_amt = 137
140 mmio_rw.write(8, add_amt)
141 read_offset_check(0, add_amt)
142 read_offset_check(12, add_amt)
143 read_offset_check(0x140, add_amt)
144
145 ##############################################################################
146 # Manifest tests
147 ##############################################################################
148
149 loopback = d.children[esi.AppID("loopback")]
150 recv = loopback.ports[esi.AppID("add")].read_port("result")
151 recv.connect()
152
153 send = loopback.ports[esi.AppID("add")].write_port("arg")
154 send.connect()
155
156 loopback_info = None
157 for mod_info in m.module_infos:
158 if mod_info.name == "LoopbackInOutAdd":
159 loopback_info = mod_info
160 break
161 assert loopback_info is not None
162 add_amt = mod_info.constants["add_amt"].value
163
164 ##############################################################################
165 # Callback tests
166 ##############################################################################
167
168 callback = d.children[esi.AppID("callback")]
169 cb_port = callback.ports[esi.AppID("cb")]
170 cb_mmio = callback.ports[esi.AppID("cmd")]
171
172 recv_data: Optional[int] = None
173
174 def my_callback(data: int) -> int:
175 nonlocal recv_data
176 recv_data = data
177 print(f"Callback received data: {data}")
178 return data + 7
179
180 cb_port.connect(my_callback)
181 cb_mmio.write(0x10, 5)
182 while recv_data is None:
183 time.sleep(0.25)
184 assert recv_data == 5
185
186 ##############################################################################
187 # Loopback add 7 tests
188 ##############################################################################
189
190 data = 10234
191 # Blocking write interface
192 send.write(data)
193 resp = recv.read()
194
195 print(f"data: {data}")
196 print(f"resp: {resp}")
197 assert resp == data + add_amt
198
199 # Non-blocking write interface
200 data = 10235
201 nb_wr_start = time.time()
202
203 # Timeout of 5 seconds
204 nb_timeout = nb_wr_start + 5
205 write_succeeded = False
206 while time.time() < nb_timeout:
207 write_succeeded = send.try_write(data)
208 if write_succeeded:
209 break
210
211 assert write_succeeded, "Non-blocking write failed"
212 resp = recv.read()
213 print(f"data: {data}")
214 print(f"resp: {resp}")
215 assert resp == data + add_amt
216
217 print("PASS")
218
219 ##############################################################################
220 # Const producer tests
221 ##############################################################################
222
223 producer_bundle = d.ports[esi.AppID("const_producer")]
224 producer = producer_bundle.read_port("data")
225 producer.connect()
226 data = producer.read()
227 producer.disconnect()
228 print(f"data: {data}")
229 assert data == 42
230
231 ##############################################################################
232 # Handshake JoinAddFunc tests
233 ##############################################################################
234
235 # Disabled test since the DC dialect flow is broken. Leaving the code here in
236 # case someone fixes it.
237
238 # a = d.ports[esi.AppID("join_a")].write_port("data")
239 # a.connect()
240 # b = d.ports[esi.AppID("join_b")].write_port("data")
241 # b.connect()
242 # x = d.ports[esi.AppID("join_x")].read_port("data")
243 # x.connect()
244
245 # a.write(15)
246 # b.write(24)
247 # xdata = x.read()
248 # print(f"join: {xdata}")
249 # assert xdata == 15 + 24
250
251 ##############################################################################
252 # StructToWindowFunc tests
253 ##############################################################################
254
255 print("Testing StructToWindowFunc...")
256 struct_to_window_bundle = d.ports[esi.AppID("struct_to_window")]
257
258 # Get the write port for sending the complete struct
259 struct_send = struct_to_window_bundle.write_port("arg")
260 struct_send.connect()
261
262 # Get the read port for receiving the windowed result
263 window_recv = struct_to_window_bundle.read_port("result")
264 window_recv.connect()
265
266 # Create test data - a struct with four 32-bit fields (as bytearrays,
267 # little-endian)
268 test_struct = {
269 "a": bytearray([0x11, 0x11, 0x11, 0x11]),
270 "b": bytearray([0x22, 0x22, 0x22, 0x22]),
271 "c": bytearray([0x33, 0x33, 0x33, 0x33]),
272 "d": bytearray([0x44, 0x44, 0x44, 0x44])
273 }
274
275 # Send the complete struct
276 struct_send.write(test_struct)
277
278 # The windowed result should arrive as two frames
279 # Frame 1 contains fields a and b
280 # Frame 2 contains fields c and d
281 # After translation, we should get back the complete struct
282 result = window_recv.read()
283
284 print(f"Sent struct: {test_struct}")
285 print(f"Received result: {result}")
286
287 # Verify the result matches the input
288 assert result["a"] == test_struct[
289 "a"], f"Field 'a' mismatch: {result['a']} != {test_struct['a']}"
290 assert result["b"] == test_struct[
291 "b"], f"Field 'b' mismatch: {result['b']} != {test_struct['b']}"
292 assert result["c"] == test_struct[
293 "c"], f"Field 'c' mismatch: {result['c']} != {test_struct['c']}"
294 assert result["d"] == test_struct[
295 "d"], f"Field 'd' mismatch: {result['d']} != {test_struct['d']}"
296
297 print("PASS: StructToWindowFunc test passed")
298
299 # Test with different values
300 test_struct2 = {
301 "a": bytearray([0xEF, 0xBE, 0xAD, 0xDE]), # 0xDEADBEEF little-endian
302 "b": bytearray([0xBE, 0xBA, 0xFE, 0xCA]), # 0xCAFEBABE little-endian
303 "c": bytearray([0x78, 0x56, 0x34, 0x12]), # 0x12345678 little-endian
304 "d":
305 bytearray([0x21, 0x43, 0x65, 0x87]) # 0x87654321 little-endian
306 }
307 struct_send.write(test_struct2)
308 result2 = window_recv.read()
309
310 print(f"Sent struct: {test_struct2}")
311 print(f"Received result: {result2}")
312
313 assert result2["a"] == test_struct2[
314 "a"], f"Field 'a' mismatch: {result2['a']} != {test_struct2['a']}"
315 assert result2["b"] == test_struct2[
316 "b"], f"Field 'b' mismatch: {result2['b']} != {test_struct2['b']}"
317 assert result2["c"] == test_struct2[
318 "c"], f"Field 'c' mismatch: {result2['c']} != {test_struct2['c']}"
319 assert result2["d"] == test_struct2[
320 "d"], f"Field 'd' mismatch: {result2['d']} != {test_struct2['d']}"
321
322 print("PASS: StructToWindowFunc test 2 passed")
323
324 struct_send.disconnect()
325 window_recv.disconnect()
326
327 ##############################################################################
328 # WindowToStructFunc tests
329 ##############################################################################
330
331 print("Testing WindowToStructFunc...")
332 window_to_struct_bundle = d.ports[esi.AppID("struct_from_window")]
333
334 # Get the write port for sending the windowed struct (two frames)
335 window_send = window_to_struct_bundle.write_port("arg")
336 window_send.connect()
337
338 # Get the read port for receiving the complete struct
339 struct_recv = window_to_struct_bundle.read_port("result")
340 struct_recv.connect()
341
342 # Create test data - a struct with four 32-bit fields (as bytearrays,
343 # little-endian). We'll send this as a windowed struct and expect to get it
344 # back as a complete struct.
345 test_window_struct = {
346 "a": bytearray([0xAA, 0xAA, 0xAA, 0xAA]),
347 "b": bytearray([0xBB, 0xBB, 0xBB, 0xBB]),
348 "c": bytearray([0xCC, 0xCC, 0xCC, 0xCC]),
349 "d": bytearray([0xDD, 0xDD, 0xDD, 0xDD])
350 }
351
352 # Send the windowed struct (the runtime will split it into two frames)
353 window_send.write(test_window_struct)
354
355 # Read the complete struct result
356 result = struct_recv.read()
357
358 print(f"Sent windowed struct: {test_window_struct}")
359 print(f"Received complete struct: {result}")
360
361 # Verify the result matches the input
362 assert result["a"] == test_window_struct[
363 "a"], f"Field 'a' mismatch: {result['a']} != {test_window_struct['a']}"
364 assert result["b"] == test_window_struct[
365 "b"], f"Field 'b' mismatch: {result['b']} != {test_window_struct['b']}"
366 assert result["c"] == test_window_struct[
367 "c"], f"Field 'c' mismatch: {result['c']} != {test_window_struct['c']}"
368 assert result["d"] == test_window_struct[
369 "d"], f"Field 'd' mismatch: {result['d']} != {test_window_struct['d']}"
370
371 print("PASS: WindowToStructFunc test passed")
372
373 # Test with different values
374 test_window_struct2 = {
375 "a": bytearray([0x01, 0x02, 0x03, 0x04]),
376 "b": bytearray([0x05, 0x06, 0x07, 0x08]),
377 "c": bytearray([0x09, 0x0A, 0x0B, 0x0C]),
378 "d": bytearray([0x0D, 0x0E, 0x0F, 0x10])
379 }
380 window_send.write(test_window_struct2)
381 result2 = struct_recv.read()
382
383 print(f"Sent windowed struct: {test_window_struct2}")
384 print(f"Received complete struct: {result2}")
385
386 assert result2["a"] == test_window_struct2[
387 "a"], f"Field 'a' mismatch: {result2['a']} != {test_window_struct2['a']}"
388 assert result2["b"] == test_window_struct2[
389 "b"], f"Field 'b' mismatch: {result2['b']} != {test_window_struct2['b']}"
390 assert result2["c"] == test_window_struct2[
391 "c"], f"Field 'c' mismatch: {result2['c']} != {test_window_struct2['c']}"
392 assert result2["d"] == test_window_struct2[
393 "d"], f"Field 'd' mismatch: {result2['d']} != {test_window_struct2['d']}"
394
395 print("PASS: WindowToStructFunc test 2 passed")
396
397 window_send.disconnect()
398 struct_recv.disconnect()
399
400
401@cosim_test(HW_DIR / "esi_test.py")
402def test_cosim_esi(conn: AcceleratorConnection) -> None:
403 run(conn)
404
405
406@cosim_test(HW_DIR / "esi_test.py")
407def test_cosim_esi_manifest_mmio(host: str, port: int) -> None:
408 os.environ["ESI_COSIM_MANIFEST_MMIO"] = "1"
409 conn = esi.connect("cosim", f"{host}:{port}")
410 run(conn)
411
412
413if __name__ == "__main__":
414 platform = sys.argv[1]
415 conn_str = sys.argv[2]
416 conn = esi.Context(esi.LogLevel.Debug).connect(platform, conn_str)
417 run(conn, platform)
static void print(TypedAttr val, llvm::raw_ostream &os)
static mlir::Operation * resolve(Context &context, mlir::SymbolRefAttr sym)
AcceleratorConnections, Accelerators, and Manifests must all share a context.
Definition Context.h:34
None test_cosim_esi(AcceleratorConnection conn)
Definition test_esi.py:402
None test_cosim_esi_manifest_mmio(str host, int port)
Definition test_esi.py:407