import torch.multiprocessing as mp POSTPROC_WORKERS = 8 # tune for optimal throughput output_queue = mp.JoinableQueue(maxsize=POSTPROC_WORKERS) def output_worker(in_q): while True: item = in_q.get() if item is None: break # signal to shut down batch_id, batch_preds = item process_output(batch_id, batch_preds) in_q.task_done() processes = [] for _ in range(POSTPROC_WORKERS): p = mp.Process(target=output_worker, args=(output_queue,)) p.start() processes.append(p) def synchronize_all(): torch.cuda.synchronize() output_queue.join() # drain queue 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 = to_cpu(output['out']) with nvtx.annotate("queue output", color="cyan"): output_queue.put((i, output_cpu)) total_time = end_time - start_time throughput = PROFILE_STEPS / total_time print(f"Throughput: {throughput:.2f} steps/sec") # cleanup for _ in range(POSTPROC_WORKERS): output_queue.put(None)