@@ -110,6 +110,7 @@ class Session:
110110 timing_scope : str = "task-model-call-wall"
111111 input_preparation_included : bool = False
112112 asset_loading_included : bool = False
113+ compile_evidence : dict [str , Any ] | None = None
113114
114115
115116def build_parser () -> argparse .ArgumentParser :
@@ -125,7 +126,9 @@ def build_parser() -> argparse.ArgumentParser:
125126 parser .add_argument ("--adapter-options-json" , default = "{}" )
126127 parser .add_argument ("--timing-contract-json" , default = "{}" )
127128 parser .add_argument ("--precision" , required = True , choices = ("fp16" , "fp32" , "bf16" ))
128- parser .add_argument ("--mode" , required = True , choices = ("hf-eager" , "pytorch-eager" ))
129+ parser .add_argument (
130+ "--mode" , required = True , choices = ("hf-eager" , "pytorch-eager" , "torch-compile" )
131+ )
129132 parser .add_argument ("--padding" , default = "longest" )
130133 parser .add_argument ("--trust-remote-code" , action = "store_true" )
131134 parser .add_argument ("--local-files-only" , action = "store_true" )
@@ -285,10 +288,25 @@ def _tensor_summary(value: Any) -> dict[str, Any]:
285288 }
286289
287290
291+ def _flatten_tensor_values (value : Any ) -> list [float ]:
292+ flattened : list [float ] = []
293+
294+ def visit (item : Any ) -> None :
295+ if isinstance (item , (list , tuple )):
296+ for child in item :
297+ visit (child )
298+ else :
299+ flattened .append (float (item ))
300+
301+ visit (value .detach ().float ().cpu ().tolist ())
302+ return flattened
303+
304+
288305def _forecast_summary (
289306 value : Any , task : str , quantile_levels : Sequence [float ] = ()
290307) -> dict [str , Any ]:
291308 summary = _tensor_summary (value )
309+ summary ["values" ] = _flatten_tensor_values (value )
292310 if task not in {"series_to_point_forecast" , "series_to_quantile_forecast" }:
293311 return summary
294312 shape = summary ["shape" ]
@@ -1352,6 +1370,9 @@ def _load_timeseries(
13521370 chronos_options = _processor_kwargs (arguments )
13531371 chronos_options .update ({"device_map" : str (device ), "dtype" : dtype })
13541372 model = ChronosBoltPipeline .from_pretrained (arguments .model , ** chronos_options )
1373+ compile_evidence = None
1374+ if arguments .mode == "torch-compile" :
1375+ compile_evidence = _compile_forward (model .model )
13551376 raw = _numeric_values (request , "past_values" )
13561377 observed = _observed_values (request , len (raw ))
13571378 context = torch .tensor ([value if mask > 0 else float ("nan" )
@@ -1367,7 +1388,7 @@ def invoke() -> Mapping[str, Any]:
13671388 quantiles = model .model .config .chronos_config ["quantiles" ]
13681389 return _forecast_summary (value , task_id , quantiles )
13691390
1370- return Session (invoke , "chronos" )
1391+ return Session (invoke , "chronos" , compile_evidence = compile_evidence )
13711392
13721393 config = transformers .AutoConfig .from_pretrained (
13731394 arguments .model , ** _processor_kwargs (arguments )
@@ -2288,18 +2309,60 @@ def _synchronize() -> None:
22882309 return
22892310
22902311
2312+ def _compile_forward (model : Any ) -> dict [str , Any ]:
2313+ import torch
2314+ from torch ._dynamo .backends .registry import lookup_backend
2315+
2316+ evidence = {"compiled_graph_count" : 0 }
2317+ inductor = lookup_backend ("inductor" )
2318+
2319+ def compile_graph (graph : Any , inputs : Any , ** options : Any ) -> Any :
2320+ compiled = inductor (graph , inputs , ** options )
2321+ evidence ["compiled_graph_count" ] += 1
2322+ return compiled
2323+
2324+ model .forward = torch .compile (
2325+ model .forward ,
2326+ backend = compile_graph ,
2327+ fullgraph = False ,
2328+ dynamic = False ,
2329+ )
2330+ evidence .update (
2331+ {
2332+ "api" : "torch.compile" ,
2333+ "target" : "model.forward" ,
2334+ "backend" : "inductor" ,
2335+ "mode" : "default" ,
2336+ "fullgraph" : False ,
2337+ "dynamic" : False ,
2338+ "applied" : True ,
2339+ }
2340+ )
2341+ return evidence
2342+
2343+
22912344def _measure (session : Session , warmup : int , iterations : int ) -> tuple [list [float ], dict [str , Any ]]:
22922345 output : Mapping [str , Any ] = {}
22932346 for _ in range (warmup ):
22942347 output = session .invoke ()
22952348 _synchronize ()
2349+ compiled_graphs = None
2350+ if session .compile_evidence is not None :
2351+ compiled_graphs = int (session .compile_evidence ["compiled_graph_count" ])
2352+ if compiled_graphs < 1 :
2353+ raise RuntimeError ("warmup did not execute a compiled graph" )
22962354 samples = []
22972355 for _ in range (iterations ):
22982356 _synchronize ()
22992357 started = time .perf_counter ()
23002358 output = session .invoke ()
23012359 _synchronize ()
23022360 samples .append ((time .perf_counter () - started ) * 1000.0 )
2361+ if (
2362+ compiled_graphs is not None
2363+ and int (session .compile_evidence ["compiled_graph_count" ]) != compiled_graphs
2364+ ):
2365+ raise RuntimeError ("model compilation occurred inside timed samples" )
23032366 return samples , dict (output )
23042367
23052368
@@ -2715,8 +2778,13 @@ def run(arguments: argparse.Namespace) -> int:
27152778 if arguments .warmup < 0 or arguments .iterations <= 0 :
27162779 raise ValueError ("warmup must be non-negative and iterations must be positive" )
27172780 expected_mode = "pytorch-eager" if arguments .adapter in PYTORCH_ADAPTERS else "hf-eager"
2718- if arguments .mode != expected_mode :
2719- raise ValueError (f"adapter { arguments .adapter } requires mode { expected_mode } " )
2781+ supported_modes = {expected_mode }
2782+ if arguments .adapter == "pytorch-timeseries" and arguments .family == "chronos_bolt" :
2783+ supported_modes .add ("torch-compile" )
2784+ if arguments .mode not in supported_modes :
2785+ raise ValueError (
2786+ f"adapter { arguments .adapter } requires one of { sorted (supported_modes )} "
2787+ )
27202788 request = flatten_config (_json_object (arguments .request_json , "--request-json" ))
27212789 options = _json_object (arguments .adapter_options_json , "--adapter-options-json" )
27222790 configured_timing = _json_object (arguments .timing_contract_json , "--timing-contract-json" )
@@ -2729,6 +2797,7 @@ def run(arguments: argparse.Namespace) -> int:
27292797 expected_timing = {name : declared [name ] for name in fields }
27302798 load_started = time .perf_counter ()
27312799 load_seconds : float | None = None
2800+ compile_evidence : dict [str , Any ] | None = None
27322801 if arguments .adapter == "upstream-elf" :
27332802 samples , output_summary , framework , timing_scope , input_included , asset_included = _run_elf (
27342803 arguments , request , options
@@ -2762,6 +2831,7 @@ def run(arguments: argparse.Namespace) -> int:
27622831 ) = _run_sana_wm (arguments , request , options )
27632832 else :
27642833 session = LOADERS [arguments .adapter ](arguments , request , options )
2834+ compile_evidence = session .compile_evidence
27652835 load_seconds = time .perf_counter () - load_started
27662836 framework = session .framework
27672837 timing_scope = session .timing_scope
@@ -2821,8 +2891,8 @@ def run(arguments: argparse.Namespace) -> int:
28212891 "precision" : arguments .precision ,
28222892 "padding" : arguments .padding ,
28232893 "experts_implementation" : None ,
2824- "compile_scope" : None ,
2825- "compile_evidence" : None ,
2894+ "compile_scope" : "model.forward" if arguments . mode == "torch-compile" else None ,
2895+ "compile_evidence" : compile_evidence ,
28262896 "timing_scope" : timing_scope ,
28272897 "input_preparation_included" : input_included ,
28282898 "asset_loading_included" : asset_included ,
@@ -2839,6 +2909,7 @@ def run(arguments: argparse.Namespace) -> int:
28392909 "input_preparation_included" : input_included ,
28402910 "asset_loading_included" : asset_included ,
28412911 "model_load_excluded" : True ,
2912+ "compile_excluded" : True ,
28422913 "warmup_excluded" : True ,
28432914 "output_materialization_included" : True ,
28442915 },
@@ -2856,6 +2927,9 @@ def run(arguments: argparse.Namespace) -> int:
28562927 "environment" : _environment (),
28572928 "finished_at" : datetime .now (timezone .utc ).isoformat (),
28582929 }
2930+ if compile_evidence is not None :
2931+ compile_evidence ["warmup_completed" ] = True
2932+ compile_evidence ["timed_callable_uses_compiled_target" ] = True
28592933 if not all (math .isfinite (float (value )) and float (value ) > 0.0 for value in samples ):
28602934 raise RuntimeError ("reference produced an invalid timing sample" )
28612935 arguments .output .parent .mkdir (parents = True , exist_ok = True )
0 commit comments