SWDEV-393541: Added new parser. Json files are served from memory.

Change-Id: I24fe7d5111ac6aead8bcf5d07960ba0a5437ae39


[ROCm/rocprofiler commit: f11ec66b1a]
Cette révision appartient à :
Giovanni LB
2023-03-29 14:35:43 -03:00
révisé par Giovanni Baraldi
Parent 1fe89b71d2
révision 4628770498
2 fichiers modifiés avec 432 ajouts et 213 suppressions
+35 -30
Voir le fichier
@@ -10,12 +10,32 @@ from struct import *
from ctypes import *
import ctypes
from copy import deepcopy
from trace_view import view_trace
from trace_view import view_trace, Readable
import sys
import glob
import numpy as np
import matplotlib.pyplot as plt
import json
from io import BytesIO
class FileBytesIO:
def __init__(self, iobytes) -> None:
self.iobytes = iobytes
self.seek = 0
def __len__(self):
return self.iobytes.getbuffer().nbytes
def read(self, length=0):
if length<=0:
return bytes(self.getbuffer())
else:
if self.seek >= len(self):
self.seek = 0
return None
response = self.iobytes.getbuffer()[self.seek:self.seek+length]
self.seek += length
return bytes(response)
COUNTERS_MAX_CAPTURES = 1<<12
@@ -166,7 +186,7 @@ def getWaves(filename, target_cu, verbose):
return waves, events
def persist(output_ui, trace_file, SIMD):
def persist(trace_file, SIMD):
trace = Path(trace_file).name
simds, waves = [], []
begin_time, end_time, timeline, instructions = [], [], [], []
@@ -278,14 +298,6 @@ def insert_waitcnt(flight_count, assembly_code):
return assembly_code
def Copy_Files(output_ui):
curpath = os.path.dirname(os.path.abspath(__file__))
outpath = output_ui+'/ui/'
os.makedirs(outpath, exist_ok=True)
os.system('cp '+curpath+'/ui/* '+outpath)
def get_delta_time(events):
try:
CUS = [[e.time for e in events if e.cu==k and e.bank==0] for k in range(16)]
@@ -295,13 +307,10 @@ def get_delta_time(events):
return 1
def draw_wave_metrics(selections, normalize):
global PIC_SAVE_FOLDER
global EVENTS
global EVENT_NAMES
#event_names = ['Busy CUs', 'Occupancy', 'Eligible waves', 'Waves waiting']
with open(os.path.join(PIC_SAVE_FOLDER,'counters.json'), 'w') as f:
f.write(json.dumps({"counters": EVENT_NAMES}))
response = Readable({"counters": EVENT_NAMES})
plt.figure(figsize=(15,3))
@@ -350,13 +359,14 @@ def draw_wave_metrics(selections, normalize):
else:
plt.ylabel('Value')
plt.subplots_adjust(left=0.05, right=1, top=1, bottom=0.07)
plt.savefig(os.path.join(PIC_SAVE_FOLDER,'timeline.png'), dpi=150)
#plt.show()
figure_bytes = BytesIO()
plt.savefig(figure_bytes, dpi=150)
return response, FileBytesIO(figure_bytes)
def draw_wave_states(selections, normalize):
global TIMELINES
global PIC_SAVE_FOLDER
plot_indices = [1, 2, 3, 4]
STATES = [['Empty', 'Idle', 'Exec', 'Wait', 'Stall'][k] for k in plot_indices]
colors = [['gray', 'orange', 'green', 'red', 'blue'][k] for k in plot_indices]
@@ -379,9 +389,6 @@ def draw_wave_states(selections, normalize):
timelines = [np.convolve(time, kernel)[kernsize//2:-kernsize//2][::trim] if len(time) > 0 else cycles*0 for time in timelines]
with open(os.path.join(PIC_SAVE_FOLDER,'counters.json'), 'w') as f:
f.write(json.dumps({"counters": STATES}))
[plt.plot(cycles, t, label='State '+s, linewidth=1.1, color=c)
for t, s, c, sel in zip(timelines, STATES, colors, selections) if sel]
@@ -393,14 +400,17 @@ def draw_wave_states(selections, normalize):
plt.ylim(-1)
plt.xlim(-maxtime//200, maxtime+maxtime//200+1)
plt.subplots_adjust(left=0.05, right=1, top=1, bottom=0.07)
plt.savefig(os.path.join(PIC_SAVE_FOLDER,'timeline.png'), dpi=150)
figure_bytes = BytesIO()
plt.savefig(figure_bytes, dpi=150)
response = Readable({"counters": STATES})
return response, FileBytesIO(figure_bytes)
def GeneratePIC(selections=[True for k in range(16)], normalize=True, bScounter=True):
if bScounter and len(EVENTS) > 0 and np.sum([len(e) for e in EVENTS]) > 32:
draw_wave_metrics(selections, normalize)
return draw_wave_metrics(selections, normalize)
else:
draw_wave_states(selections, normalize)
return draw_wave_states(selections, normalize)
if __name__ == "__main__":
@@ -410,7 +420,6 @@ if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("assembly_code", help="Path of the assembly code")
parser.add_argument("--trace_file", help="Filter for trace files", default=None, type=str)
parser.add_argument("-o", "--output_ui", help="Output Folder", default='.')
parser.add_argument("-k", "--att_kernel", help="Kernel file", type=str, default=pathenv+'/*_kernel.txt')
parser.add_argument("-p", "--ports", help="Server and websocket ports, default: 8000,18000")
parser.add_argument("--target_cu", help="Collected target CU id{0-15}", type=int, default=None)
@@ -477,7 +486,6 @@ if __name__ == "__main__":
print('Trace filenames:', filenames)
Copy_Files(args.output_ui)
DBFILES = []
global TIMELINES
global EVENTS
@@ -492,7 +500,7 @@ if __name__ == "__main__":
continue
analysed_filenames.append(name)
EVENTS.append(perfevents)
DBFILES.append( persist(args.output_ui, name, SIMD) )
DBFILES.append( persist(name, SIMD) )
for wave in SIMD:
time_acc = 0
tuples1 = wave.timeline.split('(')
@@ -512,9 +520,6 @@ if __name__ == "__main__":
TIMELINES[state[0]][time_acc:time_acc+state[1]] += 1
time_acc += state[1]
global PIC_SAVE_FOLDER
PIC_SAVE_FOLDER = os.path.abspath(os.path.join(args.output_ui, 'ui'))
if args.genasm and len(args.genasm) > 0:
flight_count = view_trace(args, 0, code, jumps, DBFILES, analysed_filenames, True, None)
+397 -183
Voir le fichier
@@ -12,21 +12,346 @@ from struct import *
from collections import defaultdict
import json
import time
#import webbrowser
import http.server
import socketserver
import socket
import asyncio
import websockets
from multiprocessing import *
from multiprocessing import Process, Manager
import numpy as np
from copy import deepcopy
from http import HTTPStatus
PORT, WebSocketPort = 8000, 18000
SP = '\u00A0'
class Readable:
def __init__(self, jsonstring) -> None:
self.jsonstr = json.dumps(jsonstring)
self.seek = 0
RS_TRACE_DEBUG = "RS_TRACE_DEBUG" in os.environ
if RS_TRACE_DEBUG:
LOG = open('./att_viewer.log', 'w')
def read(self, length=0):
if length<=0:
return self.jsonstr
else:
if self.seek >= len(self):
self.seek = 0
return None
response = self.jsonstr[self.seek:self.seek+length]
self.seek += length
return bytes(response, 'utf-8')
def __len__(self):
return len(self.jsonstr)
STACK_SIZE_LIMIT = 64
SMEM = 1
SALU = 2
VMEM = 3
FLAT = 4
LDS = 5
VALU = 6
JUMP = 7
NEXT = 8
IMMED = 9
BRANCH = 10
GETPC = 11
SETPC = 12
SWAPPC = 13
LANEIO = 14
DONT_KNOW = 100
WaveInstCategory = {
SMEM: "SMEM",
SALU: "SALU",
VMEM: "VMEM",
FLAT: "FLAT",
LDS: "LDS",
VALU: "VALU",
JUMP: "JUMP",
NEXT: "NEXT",
IMMED: "IMMED",
JUMP: "JUMP",
NEXT: "NEXT",
IMMED: "IMMED",
BRANCH: "BRANCH",
GETPC: "GETPC",
SETPC: "SETPC",
SWAPPC: "SWAPPC",
LANEIO: "LANEIO",
DONT_KNOW: "DONT_KNOW",
}
JSON_GLOBAL_DICTIONARY = {}
class RegisterWatchList:
def __init__(self, labels) -> None:
self.registers = {'v'+str(k): [[] for m in range(64)] for k in range(64)}
self.registers = {**self.registers, **{'s'+str(k): [] for k in range(64)}}
self.labels = labels
def try_translate(self, tok):
if tok[0] in ['s']:
return self.registers[self.range(tok)[0]]
elif '@' in tok:
return self.labels[tok.split('@')[0]]+1
def range(self, r):
reg = r.split(':')
if len(reg) == 1:
return reg
else:
r0 = reg[0].split('[')
return [r0[0]+str(k) for k in range(int(r0[1]), int(reg[1][:-1])+1)]
def tokenize(self, line):
return [u for u in [t.split(',')[0].strip() for t in line.split(' ')] if len(u) > 0]
def getpc(self, line, next_line):
#print('Get pc:', line)
dst = line.split(' ')[1].strip()
label_dest = next_line.split(', ')[-1].split('@')[0]
for reg in self.range(dst):
#print('Setting:', reg, label_dest, self.labels[label_dest])
self.registers[reg].append(deepcopy(self.labels[label_dest]))
def swappc(self, line, line_num):
#print('swappc pc:', line)
tokens = self.tokenize(line)
dst = tokens[1]
src = tokens[2]
#print('swap to', self.registers[self.range(src)[0]])
self.registers[self.range(dst)[0]].append(line_num+1)
popped = self.registers[self.range(src)[0]][-1]
self.registers[self.range(src)[0]] = self.registers[self.range(src)[0]][:-1]
return popped
def setpc(self, line):
#print('Set pc:', line)
src = line.split(' ')[1].strip()
#print('Going to:', self.registers[self.range(src)[0]], src)
popped = self.registers[self.range(src)[0]][-1]
self.registers[self.range(src)[0]] = self.registers[self.range(src)[0]][:-1]
return popped
def updatelane(self, line):
tokens = self.tokenize(line)
try:
#print('Lane:', tokens)
if 'v_readlane' in tokens[0]:
self.registers[tokens[1]].append(self.registers[tokens[2]][int(tokens[3])][-1])
#print('Writelane value', self.registers[tokens[2]][int(tokens[3])])
self.registers[tokens[2]][int(tokens[3])] = self.registers[tokens[2]][int(tokens[3])][:-1]
elif 'v_writelane' in tokens[0]:
self.registers[tokens[1]][int(tokens[3])].append(self.registers[tokens[2]][-1])
self.registers[tokens[2]] = self.registers[tokens[2]][-STACK_SIZE_LIMIT:]
#print('Readlane value', self.registers[tokens[2]])
except Exception as e:
#print(e, 'Could not set:', line)
pass
def try_match_swapped(insts, code, i, line):
return insts[i+1][1] == code[line][1] and insts[i][1] == code[line+1][1]
def Match(inst_value, code_value):
if code_value == inst_value:
return True
if code_value in [GETPC, SWAPPC, SETPC] and inst_value==SALU:
return True
if code_value == BRANCH and inst_value in [JUMP, NEXT]: # TODO: Maybe lets not reorder branches?
return True
return False
def get_match_lookahead(insts, code, i, line):
if try_match_swapped(insts, code, i, line):
return [i+1, i]
new_inst_order = []
allowed_insts = list(range(i, min(i+4, len(insts))))
for l in range(line, min(line+10, len(code))):
bMatch = False
for j in allowed_insts:
if Match(insts[j][1], code[l][1]):
new_inst_order.append(j)
allowed_insts.remove(j)
bMatch = True
break
if bMatch == False:
break
if len(new_inst_order):
new_inst_order += [j for j in list(range(i, max(new_inst_order)+1)) if j not in new_inst_order]
return new_inst_order
def stitch(insts, raw_code, jumps):
result, i, line, loopCount, N = [], 0, 0, defaultdict(int), len(insts)
SMEM_INST = []
VMEM_INST = []
FLAT_INST = []
NUM_SMEM = 0
NUM_VMEM = 0
NUM_FLAT = 0
mem_unroll = []
flight_count = []
labels = {}
jump_map = [0]
code = [raw_code[0]]
for c in raw_code[1:]:
c = list(c)
c[0] = c[0].split(';')[0].split('//')[0].strip()
if c[1] != 100:
code.append(c)
elif ':' in c[0]:
labels[c[0].split(':')[0]] = len(code)
jump_map.append(len(code)-1)
reverse_map = []
for k, v in enumerate(jump_map):
if v >= len(reverse_map):
reverse_map.append(k)
jumps = {jump_map[j]+1: j for j in jumps}
smem_ordering = 0
vmem_ordering = 0
max_line = 0
watchlist = RegisterWatchList(labels=labels)
num_failed_stitches = 0
MAX_FAILED_STITCHES = 128
loops = 0
maxline = 0
while i < N:
#print('L', line)
loops += 1
if line >= len(code) or loops > 100000 or num_failed_stitches >= MAX_FAILED_STITCHES:
break
maxline = max(reverse_map[line], maxline)
inst = insts[i]
as_line = code[line]
max_line = max(max_line, reverse_map[line])
matched = True
next = line+1
if as_line[1] == GETPC: # TODO: @ can put you ahead of label!
watchlist.getpc(as_line[0], code[line+1][0])
matched = inst[1] == SALU
elif as_line[1] == LANEIO:
watchlist.updatelane(as_line[0])
matched = inst[1] == VALU
elif as_line[1] == SETPC:
next = watchlist.setpc(as_line[0])
matched = inst[1] == SALU
elif as_line[1] == SWAPPC:
next = watchlist.swappc(as_line[0], line)
#print('Next:', next, code[next])
matched = inst[1] == SALU
elif inst[1] == as_line[1]:
if line in jumps:
loopCount[jumps[line]-1] += 1 # label is the previous line
num_inflight = NUM_FLAT + NUM_SMEM + NUM_VMEM
if inst[1] == SMEM or inst[1] == LDS:
smem_ordering = 1 if inst[1] == SMEM else smem_ordering
SMEM_INST.append([reverse_map[line], num_inflight])
NUM_SMEM += 1
elif inst[1] == VMEM or (inst[1] == FLAT and 'global_' in as_line[0]):
VMEM_INST.append([reverse_map[line], num_inflight])
NUM_VMEM += 1
if 'buffer_' in as_line[0]:
#watchlist.LDS_buffer_op(as_line[0])
vmem_ordering = 1
elif inst[1] == FLAT:
smem_ordering = 1
vmem_ordering = 1
FLAT_INST.append([reverse_map[line], num_inflight])
NUM_FLAT += 1
elif inst[1] == IMMED and 'waitcnt' in as_line[0]:
if 'lgkmcnt' in as_line[0]:
wait_N = int(as_line[0].split('lgkmcnt(')[1].split(')')[0])
flight_count.append([as_line[-1], num_inflight, wait_N])
if wait_N == 0:
smem_ordering = 0
if smem_ordering == 0:
offset = len(SMEM_INST)-wait_N
mem_unroll.append( [reverse_map[line], SMEM_INST[:offset]+FLAT_INST] )
SMEM_INST = SMEM_INST[offset:]
NUM_SMEM = len(SMEM_INST)
FLAT_INST = []
NUM_FLAT = 0
else:
NUM_SMEM = min(max(wait_N-NUM_FLAT, 0), NUM_SMEM)
NUM_FLAT = min(max(wait_N-NUM_SMEM, 0), NUM_FLAT)
num_inflight = NUM_FLAT + NUM_SMEM + NUM_VMEM
if 'vmcnt' in as_line[0]:
wait_N = int(as_line[0].split('vmcnt(')[1].split(')')[0])
flight_count.append([as_line[-1], num_inflight, wait_N])
if wait_N == 0:
vmem_ordering = 0
if vmem_ordering == 0:
offset = len(VMEM_INST)-wait_N
mem_unroll.append( [reverse_map[line], VMEM_INST[:offset]+FLAT_INST] )
VMEM_INST = VMEM_INST[offset:]
NUM_VMEM = len(VMEM_INST)
FLAT_INST = []
NUM_FLAT = 0
else:
NUM_VMEM = min(max(wait_N-NUM_FLAT, 0), NUM_VMEM)
NUM_FLAT = min(max(wait_N-NUM_VMEM, 0), NUM_FLAT)
elif inst[1] == JUMP and as_line[1] == BRANCH:
next = jump_map[as_line[2]]
if next is None or next == 0:
print('Jump to unknown location!', as_line)
break
elif inst[1] == NEXT and as_line[1] == BRANCH:
next = line + 1
else:
matched = False
next = line + 1
if i+1 < N and line+1 < len(code):
if try_match_swapped(insts, code, i, line):
temp = insts[i]
insts[i] = insts[i+1]
insts[i+1] = temp
next = line
elif 's_waitcnt' in as_line[0] or '_load_' in as_line[0]:
print(as_line)
break
if matched:
new_res = inst + (reverse_map[line],) # (line,)
result.append(new_res)
i += 1
num_failed_stitches = 0
else:
num_failed_stitches += 1
line = next
N = max(N, 1)
if len(result) != N:
print('Warning - Stitching rate: '+str(len(result) * 100 / N)+'% matched')
print('Leftovers:', [WaveInstCategory[insts[i+k][1]] for k in range(5) if i+k < len(insts)])
try:
print(line, code[line])
except:
pass
else:
while line < len(code):
if 's_endpgm' in code[line]:
mem_unroll.append( [reverse_map[line], SMEM_INST+VMEM_INST+FLAT_INST] )
break
line += 1
return result, loopCount, mem_unroll, flight_count, maxline
def get_ip():
@@ -43,125 +368,8 @@ def get_ip():
IPAddr = get_ip()
def debug_log(msg, last=False):
if RS_TRACE_DEBUG:
LOG.write(msg)
if last:
LOG.close()
def try_match_swapped(insts, code, i, line):
return insts[i+1][1] == code[line][1] and insts[i][1] == code[line+1][1]
def stitch(insts, code, jumps):
result, i, line, loopCount, N = [], 0, 0, defaultdict(int), len(insts)
SMEM_INST = []
VMEM_INST = []
FLAT_INST = []
NUM_SMEM = 0
NUM_VMEM = 0
NUM_FLAT = 0
mem_unroll = []
flight_count = []
smem_ordering = 0
vmem_ordering = 0
while i < N:
inst = insts[i]
if line >= len(code):
break
as_line = code[line]
if inst[1] == as_line[1]:
if line in jumps:
loopCount[line-1] += 1 # label is the previous line
matched, next = True, line + 1
num_inflight = NUM_FLAT + NUM_SMEM + NUM_VMEM
if inst[1] == 1 or inst[1] == 5: # SMEM, LDS
smem_ordering = 2 if inst[1] == 1 else smem_ordering
SMEM_INST.append([line, num_inflight])
NUM_SMEM += 1
elif inst[1] == 3 or (inst[1] == 4 and 'global_' in as_line[0]): # VMEM R/W
VMEM_INST.append([line, num_inflight])
NUM_VMEM += 1
elif inst[1] == 4: # FLAT
smem_ordering = max(smem_ordering, 1)
vmem_ordering = 1
FLAT_INST.append([line, num_inflight])
NUM_FLAT += 1
elif inst[1] == 9 and 'waitcnt' in as_line[0]:
if 'lgkmcnt' in as_line[0]:
wait_N = int(as_line[0].split('lgkmcnt(')[1].split(')')[0])
flight_count.append([as_line[-1], num_inflight, wait_N])
if wait_N == 0:
smem_ordering = 0
if smem_ordering == 0:
offset = len(SMEM_INST)-wait_N
mem_unroll.append( [line, SMEM_INST[:offset]+FLAT_INST] )
SMEM_INST = SMEM_INST[offset:]
FLAT_INST = []
NUM_FLAT = 0
NUM_SMEM = 0
else:
NUM_SMEM = min(max(wait_N-NUM_FLAT, 0), NUM_SMEM)
NUM_FLAT = min(max(wait_N-NUM_SMEM, 0), NUM_FLAT)
if 'vmcnt' in as_line[0]:
wait_N = int(as_line[0].split('vmcnt(')[1].split(')')[0])
flight_count.append([as_line[-1], num_inflight, wait_N])
if wait_N == 0:
vmem_ordering = 0
if vmem_ordering == 0:
offset = len(VMEM_INST)-wait_N
mem_unroll.append( [line, VMEM_INST[:offset]+FLAT_INST] )
VMEM_INST = VMEM_INST[offset:]
FLAT_INST = []
NUM_FLAT = 0
NUM_VMEM = 0
else:
NUM_VMEM = min(max(wait_N-NUM_FLAT, 0), NUM_VMEM)
NUM_FLAT = min(max(wait_N-NUM_VMEM, 0), NUM_FLAT)
elif inst[1] == 7 and as_line[1] == 10: # jump
matched, next = True, as_line[2]
if next is None or next == 0:
print('Jump to unknown location!', as_line)
return result, loopCount, mem_unroll, flight_count
elif inst[1] == 8 and as_line[1] == 10: # next
matched, next = True, line + 1
else:
# instructions with almost same timestamp swapped
# if i+1 < N and line+1 < len(code) and inst[0] == insts[i+1][0]:
matched = False
next = line + 1
if i+1 < N and line+1 < len(code):
if try_match_swapped(insts, code, i, line):
temp = insts[i]
insts[i] = insts[i+1]
insts[i+1] = temp
next = line
#else:
# print('Could not parse tokens:', insts[i], as_line)
if matched:
new_res = inst + (line,)
result.append(new_res)
i += 1
line = next
N = max(N, 1)
if len(result) != N:
print('Warning - Stitching rate: '+str(len(result) * 100 / N)+'% matched')
return result, loopCount, mem_unroll, flight_count
PORT, WebSocketPort = 8000, 18000
SP = '\u00A0'
def extract_tuple(content, num):
vals = content.split(',')
@@ -187,33 +395,19 @@ def get_top_n(stitched):
return top_n[:TOP_N]
def rjust_html(s, n):
s = str(s)
return SP * (n-len(s)) + s if len(s) < n else s
def rjust_html_format(msg, n1, inst, n2, n3, stall):
return str(rjust_html(msg,n1)) + str(rjust_html(inst,n2)) + str(SP*n3) + str(stall)
def wave_info(df, id):
issued_ins, mem_ins = df['issued_ins'][id], df['mem_ins'][id]
valu_ins, valu_stalls = df['valu_ins'][id], df['valu_stalls'][id]
salu_ins, salu_stalls = df['salu_ins'][id], df['salu_stalls'][id]
vmem_ins, vmem_stalls = df['vmem_ins'][id], df['vmem_stalls'][id]
smem_ins, smem_stalls = df['smem_ins'][id], df['smem_stalls'][id]
flat_ins, flat_stalls = df['flat_ins'][id], df['flat_stalls'][id]
lds_ins, lds_stalls = df['lds_ins'][id], df['lds_stalls'][id]
br_ins, br_stalls = df['br_ins'][id], df['br_stalls'][id]
return 'Issued:' + str(rjust_html(issued_ins,8)) + str(SP*2) + 'Mem:' + str(mem_ins) \
+ "-" * 26 + rjust_html_format("VALU:",6,valu_ins,8,4,valu_stalls) \
+ rjust_html_format("SALU:",6,salu_ins,8,4,salu_stalls) \
+ rjust_html_format("VMEM:",6,vmem_ins,8,4,vmem_stalls) \
+ rjust_html_format("SMEM:",6,smem_ins,8,4,smem_stalls) \
+ rjust_html_format("FLAT:",6,flat_ins,8,4,flat_stalls) \
+ rjust_html_format("LDS:",6,lds_ins,8,4,lds_stalls) \
+ rjust_html_format("BR:",6,br_ins,8,4,br_stalls)
dic = {
'Issue': df['issued_ins'][id],
'Valu': df['valu_ins'][id], 'Valu_stall': df['valu_stalls'][id],
'Salu': df['salu_ins'][id], 'Salu_stall': df['salu_stalls'][id],
'Vmem': df['vmem_ins'][id], 'Vmem_stall': df['vmem_stalls'][id],
'Smem': df['smem_ins'][id], 'Smem_stall': df['smem_stalls'][id],
'Flat': df['flat_ins'][id], 'Flat_stall': df['flat_stalls'][id],
'Lds': df['lds_ins'][id], 'Lds_stall': df['lds_stalls'][id],
'Br': df['br_ins'][id], 'Br_stall': df['br_stalls'][id],
}
dic['Issue_stall'] = int(np.sum([dic[key] for key in dic.keys() if '_STALL' in key]))
return dic
def extract_waves(waves):
@@ -248,7 +442,8 @@ def extract_waves(waves):
return result
def extract_data(df, output_ui, se_number, code, jumps):
def extract_data(df, se_number, code, jumps):
if len(df['id']) == 0 or len(df['instructions']) == 0 or len(df['timeline']) == 0:
return None
@@ -272,7 +467,7 @@ def extract_data(df, output_ui, se_number, code, jumps):
for x in df['timeline'][wave_id].split('),'):
timeline.append(extract_tuple(x, 2))
stitched, loopCount, mem_unroll, count = stitch(insts, code, jumps)
stitched, loopCount, mem_unroll, count, maxline = stitch(insts, code, jumps)
srate = len(stitched)**2 / max(len(insts), 1)
if srate <= maxgrade[df['simd'][wave_id]][df['wave_slot'][wave_id]]:
continue
@@ -290,7 +485,7 @@ def extract_data(df, output_ui, se_number, code, jumps):
"info": wave_info(df, wave_id),
"instructions": stitched,
"timeline": timeline,
"code": code,
"code": code[:maxline+16],
"waitcnt": mem_unroll
}
data_obj = {
@@ -308,20 +503,12 @@ def extract_data(df, output_ui, se_number, code, jumps):
if len(data_obj["cu_waves"]) == 0:
continue
OUT = output_ui+'/ui/se'+str(se_number)+'_sm'+str(df['simd'][wave_id])+\
'_wv'+str(df['wave_slot'][wave_id])+'.json'
#'_wv'+str(wave_id)+'.json'
with open(OUT, 'w') as f:
f.write(json.dumps(data_obj))
all_filenames.append(OUT.split('/')[-1])
OUT = 'se'+str(se_number)+'_sm'+str(df['simd'][wave_id])+'_wv'+str(df['wave_slot'][wave_id])+'.json'
JSON_GLOBAL_DICTIONARY[OUT] = Readable(data_obj)
all_filenames.append(OUT)
return flight_count, all_filenames
#def open_browser():
# time.sleep(0.1)
# webbrowser.open_new_tab('http://{0}:{1}'.format(IPAddr, PORT))
class NoCacheHTTPRequestHandler(http.server.SimpleHTTPRequestHandler):
def end_headers(self):
@@ -337,9 +524,33 @@ class NoCacheHTTPRequestHandler(http.server.SimpleHTTPRequestHandler):
global PICTURE_CALLBACK
if 'timeline.png?' in self.path:
selections = [int(s)!=0 for s in self.path.split('timeline.png?')[1]]
PICTURE_CALLBACK(selections[1:], selections[0])
#PICTURE_CALLBACK(selections[2:], selections[1], selections[0])
http.server.SimpleHTTPRequestHandler.do_GET(self)
counters_json, imagebytes = PICTURE_CALLBACK(selections[1:], selections[0])
JSON_GLOBAL_DICTIONARY['counters.json'] = counters_json
JSON_GLOBAL_DICTIONARY[self.path.split('/')[-1]] = imagebytes
if '.json' in self.path or 'timeline.png' in self.path:
try:
response_file = JSON_GLOBAL_DICTIONARY[self.path.split('/')[-1]]
#print(response_file)
except:
print('Invalid json request:', self.path)
self.send_error(HTTPStatus.NOT_FOUND, "File not found")
#print(JSON_GLOBAL_DICTIONARY.keys())
return
self.send_response(HTTPStatus.OK)
if 'timeline.png' in self.path:
self.send_header("Content-type", 'image/png')
else:
self.send_header("Content-type", 'application/json')
self.send_header("Content-Length", str(len(response_file)))
self.send_header("Last-Modified", self.date_time_string(time.time()))
self.end_headers()
self.copyfile(response_file, self.wfile)
elif self.path in ['/', '/styles.css', '/index.html', '/logo.svg']:
http.server.SimpleHTTPRequestHandler.do_GET(self)
else:
print('Invalid request:', self.path)
self.send_error(HTTPStatus.NOT_FOUND, "File not found")
class RocTCPServer(socketserver.TCPServer):
def server_bind(self):
@@ -348,10 +559,8 @@ class RocTCPServer(socketserver.TCPServer):
def run_server():
global RS_HOME
Handler = NoCacheHTTPRequestHandler
os.chdir(RS_HOME+'/ui')
os.chdir(os.path.join(os.path.dirname(os.path.abspath(__file__)),'ui'))
try:
with RocTCPServer((IPAddr, PORT), Handler) as httpd:
httpd.serve_forever()
@@ -404,17 +613,23 @@ def assign_ports(ports):
PORT, WebSocketPort = ps[0], ps[1]
def call_picture_callback(return_dict):
global PICTURE_CALLBACK
response, imagebytes = PICTURE_CALLBACK()
return_dict[0] = response
return_dict[1] = imagebytes
def view_trace(args, wait, code, jumps, dbnames, att_filenames, bReturnLoc, pic_callback):
global PICTURE_CALLBACK
PICTURE_CALLBACK = pic_callback
pic_thread = Process(target=pic_callback)
manager = Manager()
return_dict = manager.dict()
pic_thread = Process(target=call_picture_callback, args=(return_dict,))
pic_thread.start()
assert(len(dbnames) > 0)
global RS_HOME
output_ui = args.output_ui
RS_HOME = output_ui
att_filenames = [Path(f).name for f in att_filenames]
se_numbers = [int(a.split('_se')[1].split('.att')[0]) for a in att_filenames]
flight_count = []
@@ -424,7 +639,7 @@ def view_trace(args, wait, code, jumps, dbnames, att_filenames, bReturnLoc, pic_
if len(dbname['id']) == 0:
continue
count, wv_filenames = extract_data(dbname, output_ui, se_number, code, jumps)
count, wv_filenames = extract_data(dbname, se_number, code, jumps)
if count is not None:
flight_count.append(count)
@@ -452,8 +667,7 @@ def view_trace(args, wait, code, jumps, dbnames, att_filenames, bReturnLoc, pic_
simd_wave_filenames[key] = wv_dict
with open(output_ui+'/ui/filenames.json', 'w') as f:
f.write(json.dumps({"filenames": simd_wave_filenames}))
JSON_GLOBAL_DICTIONARY['filenames.json'] = Readable({"filenames": simd_wave_filenames})
if args.ports:
assign_ports(args.ports)
@@ -461,11 +675,11 @@ def view_trace(args, wait, code, jumps, dbnames, att_filenames, bReturnLoc, pic_
if wait == 0:
try:
PROCS = [Process(target=run_server),
#Process(target=open_browser),
Process(target=run_websocket)]
PROCS = [Process(target=run_server), Process(target=run_websocket)]
if pic_thread is not None:
pic_thread.join()
JSON_GLOBAL_DICTIONARY['counters.json'] = return_dict[0]
JSON_GLOBAL_DICTIONARY['timeline.png'] = return_dict[1]
for p in PROCS:
p.start()