"""Fail-closed official FLA strict-config setup for one isolated process. No model/prompt/head changes. The guard prevents FLA STRICT's ordinary missing key fallback and records actual configured calls. This is a diagnostic module, not an installed change to the published wrapper or dependency environment. """ import hashlib,importlib,json,os,sys from pathlib import Path def sha(p):return hashlib.sha256(Path(p).read_bytes()).hexdigest() def serialized(key):return json.dumps(key,separators=(',',':'),sort_keys=True) def config_fields(config): if isinstance(config,dict):return {k:config.get(k) for k in ['kwargs','num_warps','num_stages','num_ctas','maxnreg','ir_override']} return {k:getattr(config,k,None) for k in ['kwargs','num_warps','num_stages','num_ctas','maxnreg','ir_override']} def validate_profile(path,expected_sha): path=Path(path).resolve() if sha(path)!=expected_sha:raise ValueError('Profile hash changed') profile=json.loads(path.read_text()) if profile['format']!='decision-fla-l2norm-profile-v1' or profile.get('model_family')!='Qwen/Qwen3.5-4B' or profile['cache_mode']!='strict':raise ValueError('Profile format/mode unsupported') if len(profile['files'])!=1 or profile['files'][0]['file']!='l2norm_fwd_kernel.json':raise ValueError('Unexpected profile file set') f=path.parent/'l2norm_fwd_kernel.json' if sha(f)!=profile['files'][0]['sha256']:raise ValueError('Explicit kernel config changed') data=json.loads(f.read_text());entries={} if data.get('default_config') is not None:raise ValueError('Implicit fallback defaults forbidden') for h,item in data['autotune_entries'].items(): key=item['autotune_key'];encoded=serialized(key) if hashlib.md5(encoded.encode()).hexdigest()!=h or encoded in entries:raise ValueError('Invalid/duplicate key') if len(key)!=5 or key[0]!=128 or type(key[1]) is not int or not 1<=key[1]<=64 or key[2:]!=['torch.bfloat16','torch.bfloat16','torch.float32']:raise ValueError('Unsupported numerical key') c=item['config'] if c['kwargs'].keys()!={'BT'} or c['kwargs']['BT'] not in [8,16,32,64] or c['num_warps'] not in [1,2,4,8,16] or c['num_stages']!=3 or c['num_ctas']!=1 or any(c.get(x) is not None for x in ['maxnreg','pre_hook','ir_override']):raise ValueError('Unexpected launch configuration') entries[encoded]=c if {json.loads(k)[1] for k in entries}!=set(range(1,65)):raise ValueError('Incomplete legal NB coverage') return path,profile,entries def attach_guard(kernel,cache_module,entries,telemetry): original=kernel.run def guarded(*args,**kwargs): if cache_module.FLA_CACHE_MODE is not cache_module.FlaCacheMode.STRICT:raise RuntimeError('FLA strict mode changed') key=cache_module.AutotuneKey.build(kernel.arg_names,kernel.keys,args,kwargs);encoded=serialized(list(key.autotune_key)) if encoded not in entries:raise RuntimeError('Uncontracted FLA l2norm key: '+encoded) expected=entries[encoded];loaded=cache_module.load_cached_config(kernel.kernel_name,key) if config_fields(loaded)!=config_fields(expected):raise RuntimeError('FLA exact config lookup mismatch') if key.autotune_key in kernel.cache and config_fields(kernel.cache[key.autotune_key])!=config_fields(expected):raise RuntimeError('A conflicting in-process kernel cache exists') # Explicit official configuration load guarantees that the following original # run finds this exact cache entry and cannot perform timing-based autotune. kernel.maybe_load_cached_config(key) if key.autotune_key not in kernel.cache or config_fields(kernel.cache[key.autotune_key])!=config_fields(expected):raise RuntimeError('Official strict config did not load') result=original(*args,**kwargs) if config_fields(kernel.cache[key.autotune_key])!=config_fields(expected):raise RuntimeError('Kernel config changed during call') telemetry['calls']+=1;telemetry['keys'][encoded]=telemetry['keys'].get(encoded,0)+1 return result kernel.run=guarded return original def install(profile_path,expected_sha): path,profile,entries=validate_profile(profile_path,expected_sha) if any(n=='fla' or n.startswith('fla.') for n in sys.modules):raise RuntimeError('Install profile before importing FLA; use a fresh isolated process') for name,wanted in {'FLA_CACHE_MODE':'strict','FLA_CONFIG_DIR':str(path.parent)}.items(): actual=os.environ.get(name) if actual is not None and actual!=wanted:raise RuntimeError('Conflicting '+name) os.environ[name]=wanted import torch,triton,fla actual={'torch':str(torch.__version__),'hip':torch.version.hip,'triton':triton.__version__,'fla':fla.__version__} for name,value in actual.items(): if value!=profile['runtime'][name]:raise RuntimeError('Runtime mismatch: '+name) if not torch.cuda.is_available():raise RuntimeError('Profile is only qualified for the specified ROCm GPU') arch=torch.cuda.get_device_properties(0).gcnArchName.split(':')[0] if arch!=profile['runtime']['gpu_arch']:raise RuntimeError('Unsupported GPU architecture '+arch) module=importlib.import_module('fla.modules.l2norm');cache_module=importlib.import_module('fla.ops.utils.cache');root=Path(fla.__file__).parent for name,value in profile['fla_source_sha256'].items(): if sha(root/name)!=value:raise RuntimeError('Pinned FLA source changed: '+name) kernel=module.l2norm_fwd_kernel if kernel.kernel_name!='l2norm_fwd_kernel' or kernel.keys!=['D','NB'] or kernel.cache:raise RuntimeError('Kernel identity or fresh-cache precondition failed') telemetry={'profile_sha256':expected_sha,'status':'installed','calls':0,'keys':{},'strict_guard':True,'unknown_keys':'raise','autotune_fallback_permitted':False,'runtime':actual,'gpu_arch':arch,'process_scope':'one explicitly profiled Nox model; other model loading in this process is not supported'} attach_guard(kernel,cache_module,entries,telemetry) # The validated inference path only uses the vectorized D128 forward kernel. # Other dimensions/backward must not silently enter a different autotuner. for name in ['l2norm_fwd_kernel1','l2norm_bwd_kernel','l2norm_bwd_kernel1']: other=getattr(module,name) def reject(*args,_name=name,**kwargs):raise RuntimeError('Uncontracted normalization kernel: '+_name) other.run=reject return telemetry