Skip to content

Commit 8ab42ae

Browse files
committed
Fix perfetto import formatting
1 parent c057365 commit 8ab42ae

3 files changed

Lines changed: 25 additions & 17 deletions

File tree

transformer_nuggets/utils/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -19,7 +19,6 @@
1919
from transformer_nuggets.utils.memory_viz import generate_memory_comparison_html
2020
from transformer_nuggets.utils.perfetto import (
2121
TraceFormat,
22-
chrome_trace_to_track_event_trace,
2322
default_trace_path,
2423
default_track_event_path,
2524
perfetto_trace_path,
@@ -29,4 +28,5 @@
2928
write_trace,
3029
write_track_event_trace,
3130
)
31+
from transformer_nuggets.utils.track_event import chrome_trace_to_track_event_trace
3232
# from transformer_nuggets.utils.model_extraction import extract_attention_data

transformer_nuggets/utils/perfetto.py

Lines changed: 7 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -24,10 +24,15 @@
2424
from collections import defaultdict
2525
from contextlib import contextmanager
2626
from dataclasses import dataclass
27-
from pathlib import Path
28-
from typing import Any, Literal
2927
from collections.abc import Iterator
28+
from pathlib import Path
3029
from re import Pattern
30+
from typing import Any, Literal
31+
32+
from transformer_nuggets.utils.track_event import (
33+
default_track_event_path,
34+
write_track_event_trace,
35+
)
3136

3237

3338
TraceFormat = Literal["chrome_json", "track_event"]
@@ -426,13 +431,6 @@ def _reassign_sort_indices(
426431
)
427432

428433

429-
from transformer_nuggets.utils.track_event import (
430-
chrome_trace_to_track_event_trace,
431-
default_track_event_path,
432-
write_track_event_trace,
433-
)
434-
435-
436434
def perfetto_trace_path(
437435
path: str | Path,
438436
*,

transformer_nuggets/utils/track_event.py

Lines changed: 17 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -527,10 +527,14 @@ def _paired_flow_ids_by_slice(
527527
flow_ids_by_slice[source_slice.index].add(flow_id)
528528

529529
for destination in destinations:
530-
destination_slice = min(
531-
kernel_slices,
532-
key=lambda slc: abs(slc.ts_us - destination.ts_us),
533-
) if kernel_slices else None
530+
destination_slice = (
531+
min(
532+
kernel_slices,
533+
key=lambda slc: abs(slc.ts_us - destination.ts_us),
534+
)
535+
if kernel_slices
536+
else None
537+
)
534538
if destination_slice is None:
535539
destination_slice = _smallest_containing_slice(slices, destination)
536540
if destination_slice is None:
@@ -814,9 +818,15 @@ def _emit_duration_markers(
814818
add_packet = builder.add_packet
815819
slice_begin = protos.TrackEvent.TYPE_SLICE_BEGIN
816820
slice_end = protos.TrackEvent.TYPE_SLICE_END
817-
for ts_ns, _begin_order, _duration_key, track_uuid, _slice_index, is_begin, slc in (
818-
_duration_markers(trace, track_ids)
819-
):
821+
for (
822+
ts_ns,
823+
_begin_order,
824+
_duration_key,
825+
track_uuid,
826+
_slice_index,
827+
is_begin,
828+
slc,
829+
) in _duration_markers(trace, track_ids):
820830
packet = add_packet()
821831
packet.timestamp = ts_ns
822832
packet.trusted_packet_sequence_id = trusted_packet_sequence_id

0 commit comments

Comments
 (0)