Add EPD disaggregated encode tracing (#25994)
This commit is contained in:
@@ -1684,10 +1684,10 @@ class MMReceiverBase(ABC):
|
|||||||
finally:
|
finally:
|
||||||
recv_socket.close()
|
recv_socket.close()
|
||||||
|
|
||||||
def send_encode_request(self, obj):
|
def send_encode_request(self, obj, time_stats_json=None):
|
||||||
self._send_encode_request(obj)
|
self._send_encode_request(obj, time_stats_json=time_stats_json)
|
||||||
|
|
||||||
def _send_encode_request(self, obj):
|
def _send_encode_request(self, obj, time_stats_json=None):
|
||||||
mm_data = self._extract_url_data(obj)
|
mm_data = self._extract_url_data(obj)
|
||||||
if obj.rid is None:
|
if obj.rid is None:
|
||||||
obj.rid = uuid.uuid4().hex
|
obj.rid = uuid.uuid4().hex
|
||||||
@@ -1730,6 +1730,7 @@ class MMReceiverBase(ABC):
|
|||||||
num_items_assigned,
|
num_items_assigned,
|
||||||
None,
|
None,
|
||||||
encode_urls,
|
encode_urls,
|
||||||
|
time_stats_json,
|
||||||
),
|
),
|
||||||
daemon=True,
|
daemon=True,
|
||||||
)
|
)
|
||||||
@@ -1839,6 +1840,7 @@ class MMReceiverBase(ABC):
|
|||||||
num_items_assigned,
|
num_items_assigned,
|
||||||
embedding_port,
|
embedding_port,
|
||||||
encode_urls=None,
|
encode_urls=None,
|
||||||
|
time_stats_json=None,
|
||||||
):
|
):
|
||||||
try:
|
try:
|
||||||
asyncio.run(
|
asyncio.run(
|
||||||
@@ -1850,6 +1852,7 @@ class MMReceiverBase(ABC):
|
|||||||
endpoint_send=None,
|
endpoint_send=None,
|
||||||
num_items_assigned=num_items_assigned,
|
num_items_assigned=num_items_assigned,
|
||||||
encode_urls=encode_urls,
|
encode_urls=encode_urls,
|
||||||
|
time_stats_json=time_stats_json,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -2069,6 +2072,7 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
endpoint_send,
|
endpoint_send,
|
||||||
num_items_assigned=None,
|
num_items_assigned=None,
|
||||||
encode_urls=None,
|
encode_urls=None,
|
||||||
|
time_stats_json=None,
|
||||||
):
|
):
|
||||||
if len(mm_data) == 0:
|
if len(mm_data) == 0:
|
||||||
return
|
return
|
||||||
@@ -2121,6 +2125,7 @@ class MMReceiverHTTP(MMReceiverBase):
|
|||||||
"modality": modality.name, # convert enum to string for json serialization
|
"modality": modality.name, # convert enum to string for json serialization
|
||||||
"prefill_host": self.host,
|
"prefill_host": self.host,
|
||||||
"embedding_port": embedding_port,
|
"embedding_port": embedding_port,
|
||||||
|
"time_stats_json": time_stats_json,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
cum_idx += 1
|
cum_idx += 1
|
||||||
|
|||||||
@@ -48,6 +48,11 @@ from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem
|
|||||||
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
||||||
from sglang.srt.model_loader import get_model
|
from sglang.srt.model_loader import get_model
|
||||||
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
|
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
|
||||||
|
from sglang.srt.observability.req_time_stats import EncoderReqTimeStats
|
||||||
|
from sglang.srt.observability.trace import (
|
||||||
|
process_tracing_init,
|
||||||
|
trace_set_thread_info,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -2447,6 +2452,10 @@ async def _dp_worker_encode_and_send(
|
|||||||
# Mooncake returns metadata for main to forward; zmq inlines the send.
|
# Mooncake returns metadata for main to forward; zmq inlines the send.
|
||||||
# Soft errors raise MMError so the dispatcher route maps them to HTTP.
|
# Soft errors raise MMError so the dispatcher route maps them to HTTP.
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
|
time_stats_json = request.pop("time_stats_json", None)
|
||||||
|
time_stats = EncoderReqTimeStats()
|
||||||
|
if time_stats_json:
|
||||||
|
time_stats.decode_json(time_stats_json)
|
||||||
request["enter_time"] = time.time()
|
request["enter_time"] = time.time()
|
||||||
modality = Modality.from_str(request["modality"])
|
modality = Modality.from_str(request["modality"])
|
||||||
backend = enc.server_args.encoder_transfer_backend
|
backend = enc.server_args.encoder_transfer_backend
|
||||||
@@ -2461,14 +2470,20 @@ async def _dp_worker_encode_and_send(
|
|||||||
code=HTTPStatus.BAD_REQUEST,
|
code=HTTPStatus.BAD_REQUEST,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
time_stats.set_mm_encode_start_time()
|
||||||
encode_coro = (
|
encode_coro = (
|
||||||
sched.submit(request)
|
sched.submit(request)
|
||||||
if sched is not None and modality in _BATCHABLE_MODALITIES
|
if sched is not None and modality in _BATCHABLE_MODALITIES
|
||||||
else enc.encode_request(request, modality)
|
else enc.encode_request(request, modality)
|
||||||
)
|
)
|
||||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = await encode_coro
|
try:
|
||||||
|
nbytes, embedding_len, embedding_dim, error_msg, error_code = await encode_coro
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
time_stats.trace_ctx.abort(abort_info={"reason": "encoder batch timed out"})
|
||||||
|
raise
|
||||||
|
|
||||||
if error_msg:
|
if error_msg:
|
||||||
|
time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
|
||||||
# zmq backends still forward an error EmbeddingData to P so it
|
# zmq backends still forward an error EmbeddingData to P so it
|
||||||
# doesn't block; send failures here are swallowed.
|
# doesn't block; send failures here are swallowed.
|
||||||
try:
|
try:
|
||||||
@@ -2490,6 +2505,8 @@ async def _dp_worker_encode_and_send(
|
|||||||
enc.embedding_to_send.pop(req_id, None)
|
enc.embedding_to_send.pop(req_id, None)
|
||||||
raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR)
|
raise MMError(error_msg, code=error_code or HTTPStatus.INTERNAL_SERVER_ERROR)
|
||||||
|
|
||||||
|
time_stats.set_mm_encode_end_time()
|
||||||
|
|
||||||
if backend == "mooncake":
|
if backend == "mooncake":
|
||||||
request.pop("mm_items", None)
|
request.pop("mm_items", None)
|
||||||
request.update(
|
request.update(
|
||||||
@@ -3275,6 +3292,13 @@ def launch_server(server_args: ServerArgs):
|
|||||||
dist_init_method = NetworkAddress(
|
dist_init_method = NetworkAddress(
|
||||||
server_args.host or "127.0.0.1", port_args.nccl_port
|
server_args.host or "127.0.0.1", port_args.nccl_port
|
||||||
).to_tcp()
|
).to_tcp()
|
||||||
|
if server_args.enable_trace:
|
||||||
|
process_tracing_init(
|
||||||
|
server_args.otlp_traces_endpoint,
|
||||||
|
"sglang",
|
||||||
|
trace_modules=server_args.trace_modules,
|
||||||
|
)
|
||||||
|
trace_set_thread_info("Encoder")
|
||||||
for rank in range(1, server_args.tp_size):
|
for rank in range(1, server_args.tp_size):
|
||||||
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
|
||||||
send_sockets.append(
|
send_sockets.append(
|
||||||
@@ -3414,7 +3438,12 @@ async def get_condition(rid):
|
|||||||
async def handle_encode_request(request: dict):
|
async def handle_encode_request(request: dict):
|
||||||
req_id = request["req_id"]
|
req_id = request["req_id"]
|
||||||
start_time = time.monotonic()
|
start_time = time.monotonic()
|
||||||
|
time_stats_json = request.pop("time_stats_json", None)
|
||||||
|
time_stats = EncoderReqTimeStats()
|
||||||
if dp_dispatcher is not None:
|
if dp_dispatcher is not None:
|
||||||
|
if time_stats_json:
|
||||||
|
request = dict(request)
|
||||||
|
request["time_stats_json"] = time_stats_json
|
||||||
try:
|
try:
|
||||||
result = await dp_dispatcher.dispatch(request)
|
result = await dp_dispatcher.dispatch(request)
|
||||||
except MMError as e:
|
except MMError as e:
|
||||||
@@ -3497,12 +3526,19 @@ async def handle_encode_request(request: dict):
|
|||||||
async with encoder.encode_dispatch_lock:
|
async with encoder.encode_dispatch_lock:
|
||||||
request.update({"enter_time": time.time()})
|
request.update({"enter_time": time.time()})
|
||||||
modality = Modality.from_str(request["modality"])
|
modality = Modality.from_str(request["modality"])
|
||||||
|
if time_stats_json:
|
||||||
|
time_stats.decode_json(time_stats_json)
|
||||||
|
|
||||||
|
time_stats.set_mm_encode_start_time()
|
||||||
if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
|
if encoder_scheduler is not None and modality in _BATCHABLE_MODALITIES:
|
||||||
try:
|
try:
|
||||||
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
nbytes, embedding_len, embedding_dim, error_msg, error_code = (
|
||||||
await encoder_scheduler.submit(request)
|
await encoder_scheduler.submit(request)
|
||||||
)
|
)
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
|
time_stats.trace_ctx.abort(
|
||||||
|
abort_info={"reason": "encoder batch timed out"}
|
||||||
|
)
|
||||||
return ORJSONResponse(
|
return ORJSONResponse(
|
||||||
status_code=HTTPStatus.GATEWAY_TIMEOUT,
|
status_code=HTTPStatus.GATEWAY_TIMEOUT,
|
||||||
content={
|
content={
|
||||||
@@ -3518,6 +3554,11 @@ async def handle_encode_request(request: dict):
|
|||||||
await encoder.encode_request(request, modality)
|
await encoder.encode_request(request, modality)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if error_msg:
|
||||||
|
time_stats.trace_ctx.abort(abort_info={"reason": error_msg})
|
||||||
|
else:
|
||||||
|
time_stats.set_mm_encode_end_time()
|
||||||
|
|
||||||
if error_msg:
|
if error_msg:
|
||||||
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
|
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
|
||||||
if request["embedding_port"] is None:
|
if request["embedding_port"] is None:
|
||||||
|
|||||||
@@ -2860,7 +2860,15 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
"zmq_to_scheduler",
|
"zmq_to_scheduler",
|
||||||
"mooncake",
|
"mooncake",
|
||||||
]:
|
]:
|
||||||
self.mm_receiver.send_encode_request(obj)
|
time_stats_json = None
|
||||||
|
if self.server_args.enable_trace:
|
||||||
|
state = self.rid_to_state.get(obj.rid)
|
||||||
|
if state is not None:
|
||||||
|
time_stats_json = state.time_stats.encode_json()
|
||||||
|
|
||||||
|
self.mm_receiver.send_encode_request(
|
||||||
|
obj, time_stats_json=time_stats_json
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
obj.need_wait_for_mm_inputs = False
|
obj.need_wait_for_mm_inputs = False
|
||||||
|
|
||||||
|
|||||||
@@ -145,6 +145,12 @@ class RequestStage:
|
|||||||
metrics_is_observed=True,
|
metrics_is_observed=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# EPD disaggregation Encode process
|
||||||
|
MM_ENCODE = RequestStageConfig(
|
||||||
|
"mm_encode",
|
||||||
|
level=1,
|
||||||
|
)
|
||||||
|
|
||||||
# disaggregation prefill
|
# disaggregation prefill
|
||||||
PREFILL_PREPARE = RequestStageConfig(
|
PREFILL_PREPARE = RequestStageConfig(
|
||||||
"prefill_prepare",
|
"prefill_prepare",
|
||||||
@@ -306,14 +312,34 @@ class ReqTimeStatsBase:
|
|||||||
def __getstate__(self) -> object:
|
def __getstate__(self) -> object:
|
||||||
# The object is propagated to other processes via serialization and deserialization methods,
|
# The object is propagated to other processes via serialization and deserialization methods,
|
||||||
# requiring the metric collector to be reconfigured.
|
# requiring the metric collector to be reconfigured.
|
||||||
|
trace_ctx_state = (
|
||||||
|
self.trace_ctx.__getstate__()
|
||||||
|
if self.trace_ctx.tracing_enable
|
||||||
|
else {"tracing_enable": False}
|
||||||
|
)
|
||||||
return {
|
return {
|
||||||
"disagg_mode": self.disagg_mode,
|
"disagg_mode": self.disagg_mode.value if self.disagg_mode else None,
|
||||||
"enable_metrics": False,
|
"enable_metrics": False,
|
||||||
"trace_ctx": self.trace_ctx,
|
"trace_ctx": trace_ctx_state,
|
||||||
"diff_realtime_monotonic": global_diff_realtime_monotonic,
|
"diff_realtime_monotonic": global_diff_realtime_monotonic,
|
||||||
}
|
}
|
||||||
|
|
||||||
def __setstate__(self, state: object):
|
def __setstate__(self, state: object):
|
||||||
|
# Reconstruct disagg_mode from string value if needed
|
||||||
|
disagg_mode_val = state.get("disagg_mode")
|
||||||
|
if isinstance(disagg_mode_val, str):
|
||||||
|
state["disagg_mode"] = DisaggregationMode(disagg_mode_val)
|
||||||
|
|
||||||
|
# Reconstruct trace_ctx from serialized dict if needed
|
||||||
|
trace_ctx_state = state.get("trace_ctx")
|
||||||
|
if isinstance(trace_ctx_state, dict):
|
||||||
|
if trace_ctx_state.get("tracing_enable"):
|
||||||
|
trace_ctx = object.__new__(TraceReqContext)
|
||||||
|
trace_ctx.__setstate__(trace_ctx_state)
|
||||||
|
state["trace_ctx"] = trace_ctx
|
||||||
|
else:
|
||||||
|
state["trace_ctx"] = TraceNullContext()
|
||||||
|
|
||||||
for key in state.keys():
|
for key in state.keys():
|
||||||
if key.endswith("time"):
|
if key.endswith("time"):
|
||||||
state[key] = convert_time_cross_thread(
|
state[key] = convert_time_cross_thread(
|
||||||
@@ -323,6 +349,12 @@ class ReqTimeStatsBase:
|
|||||||
)
|
)
|
||||||
self.__dict__.update(state)
|
self.__dict__.update(state)
|
||||||
|
|
||||||
|
def encode_json(self) -> Dict[str, Any]:
|
||||||
|
return self.__getstate__()
|
||||||
|
|
||||||
|
def decode_json(self, state: Dict[str, Any]):
|
||||||
|
self.__setstate__(state)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class APIServerReqTimeStats(ReqTimeStatsBase):
|
class APIServerReqTimeStats(ReqTimeStatsBase):
|
||||||
@@ -353,6 +385,13 @@ class APIServerReqTimeStats(ReqTimeStatsBase):
|
|||||||
|
|
||||||
if self.trace_ctx.tracing_enable:
|
if self.trace_ctx.tracing_enable:
|
||||||
self.trace_ctx.trace_req_start(convert_time_to_realtime_ns(ts))
|
self.trace_ctx.trace_req_start(convert_time_to_realtime_ns(ts))
|
||||||
|
# Start tokenize span early so that EPD encode dispatch can capture
|
||||||
|
# it as the predecessor span context when serializing trace_ctx.
|
||||||
|
self.trace_ctx.trace_slice_start(
|
||||||
|
RequestStage.TOKENIZE.stage_name,
|
||||||
|
RequestStage.TOKENIZE.level,
|
||||||
|
convert_time_to_realtime_ns(ts),
|
||||||
|
)
|
||||||
|
|
||||||
def set_finished_time(self, ts=None):
|
def set_finished_time(self, ts=None):
|
||||||
ts = ts or time.perf_counter()
|
ts = ts or time.perf_counter()
|
||||||
@@ -374,8 +413,13 @@ class APIServerReqTimeStats(ReqTimeStatsBase):
|
|||||||
ts = ts or time.perf_counter()
|
ts = ts or time.perf_counter()
|
||||||
self.tokenize_finish_time = ts
|
self.tokenize_finish_time = ts
|
||||||
|
|
||||||
stage = RequestStage.TOKENIZE
|
# tokenize span was started in set_created_time(); end it here.
|
||||||
self.trace_slice(stage, self.created_time, ts)
|
if self.trace_ctx.tracing_enable:
|
||||||
|
self.trace_ctx.trace_slice_end(
|
||||||
|
RequestStage.TOKENIZE.stage_name,
|
||||||
|
RequestStage.TOKENIZE.level,
|
||||||
|
convert_time_to_realtime_ns(ts),
|
||||||
|
)
|
||||||
|
|
||||||
def set_api_server_dispatch_time(self, ts=None):
|
def set_api_server_dispatch_time(self, ts=None):
|
||||||
ts = ts or time.perf_counter()
|
ts = ts or time.perf_counter()
|
||||||
@@ -1141,6 +1185,34 @@ class SchedulerReqTimeStats(ReqTimeStatsBase):
|
|||||||
return f"{convert_time_to_realtime(perf_counter_time):.3f}"
|
return f"{convert_time_to_realtime(perf_counter_time):.3f}"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class EncoderReqTimeStats(ReqTimeStatsBase):
|
||||||
|
mm_encode_start_time: float = 0.0
|
||||||
|
mm_encode_end_time: float = 0.0
|
||||||
|
|
||||||
|
def set_mm_encode_start_time(self, ts=None):
|
||||||
|
ts = ts or time.perf_counter()
|
||||||
|
self.mm_encode_start_time = ts
|
||||||
|
if self.trace_ctx.tracing_enable:
|
||||||
|
self.trace_ctx.rebuild_thread_context()
|
||||||
|
self.trace_ctx.trace_slice_start(
|
||||||
|
RequestStage.MM_ENCODE.stage_name,
|
||||||
|
RequestStage.MM_ENCODE.level,
|
||||||
|
convert_time_to_realtime_ns(ts),
|
||||||
|
)
|
||||||
|
|
||||||
|
def set_mm_encode_end_time(self, ts=None):
|
||||||
|
ts = ts or time.perf_counter()
|
||||||
|
self.mm_encode_end_time = ts
|
||||||
|
if self.trace_ctx.tracing_enable:
|
||||||
|
self.trace_ctx.trace_slice_end(
|
||||||
|
RequestStage.MM_ENCODE.stage_name,
|
||||||
|
RequestStage.MM_ENCODE.level,
|
||||||
|
convert_time_to_realtime_ns(ts),
|
||||||
|
thread_finish_flag=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def set_schedule_time_batch(batch: ScheduleBatch):
|
def set_schedule_time_batch(batch: ScheduleBatch):
|
||||||
# only for tracing
|
# only for tracing
|
||||||
if not get_global_tracing_enabled():
|
if not get_global_tracing_enabled():
|
||||||
|
|||||||
Reference in New Issue
Block a user