Decision-1.0-Nox-4B / code /profile_guard.py
Xunzhuo's picture
Release measured v1.2 update
d330a0b verified
Raw History Blame
6.18 kB
"""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