shape = (BATCH_SIZE, N_CLASSES, IMG_SIZE, IMG_SIZE) buffer_pool = [torch.empty(shape).share_memory_() for _ in range(POSTPROC_WORKERS)] buf_queue = mp.Queue() for i in range(POSTPROC_WORKERS): buf_queue.put(i) def output_worker(buffer_pool, in_q, buf_q): while True: item = in_q.get() if item is None: break # signal to shut down batch_id, buf_id = item process_output(batch_id, buffer_pool[buf_id]) buf_q.put(buf_id) in_q.task_done() processes = [] for _ in range(POSTPROC_WORKERS): p = mp.Process(target=output_worker, args=(buffer_pool,output_queue,buf_queue)) p.start() processes.append(p) def to_cpu(output): buf_id = buf_queue.get() output_cpu = buffer_pool[buf_id] output_cpu.copy_(output) return output_cpu, buf_id with torch.inference_mode(): for i in range(TOTAL_STEPS): if i == WARMUP_STEPS: synchronize_all() start_time = time.perf_counter() profiler.start() elif i == WARMUP_STEPS + PROFILE_STEPS: synchronize_all() profiler.stop() end_time = time.perf_counter() with nvtx.annotate(f"Batch {i}", color="blue"): with nvtx.annotate("get batch", color="red"): batch = next(data_iter) with nvtx.annotate("compute", color="green"): output = model(batch) with nvtx.annotate("copy to CPU", color="yellow"): output_cpu, buf_id = to_cpu(output['out']) with nvtx.annotate("queue output", color="cyan"): output_queue.put((i, buf_id))