litert-models / conversion /build_trainable_fire_multitask.py
unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified
Raw History Blame Contribute Delete
8.89 kB
"""DINOv3 pilot: mutable output adapters for all four original tasks.
Frozen DINOv3/DPT; preserves original logits at initialization. Partial labels
are explicit: missing supervision is never silently converted to a negative.
"""
import os
os.environ.update(TF_CPP_MIN_LOG_LEVEL='2',TF_NUM_INTRAOP_THREADS='6',TF_NUM_INTEROP_THREADS='2')
import json,hashlib
from pathlib import Path
import tensorflow as tf,numpy as np
from merge_training_graph import merge,dynamic_resize_backbone
root=Path('/workspace/vds-litert-batch2');key='fireviewer_dinov3_pointing_pilot_v1';out=root/'learning'/key;out.mkdir(exist_ok=True)
S=448
names=['abstention_logits','point_logits','presence_logits','segmentation_logits']
shapes=[[1],[1,1,S,S],[1,2],[1,1,S,S]]
features=[tf.TensorSpec(shape,tf.float32,name='features_'+name) for name,shape in zip(names,shapes)]
class Head(tf.Module):
def __init__(self):
super().__init__()
self.scale=tf.Variable(tf.ones([4]),name='scale');self.bias=tf.Variable(tf.zeros([4]),name='bias')
self.pw=tf.Variable(tf.eye(2),name='presence_weight');self.pb=tf.Variable(tf.zeros([2]),name='presence_bias')
self.parameters=[self.scale,self.bias,self.pw,self.pb]
self.momentum=[tf.Variable(tf.zeros_like(v),name='momentum_'+v.name.split(':')[0]) for v in self.parameters]
self.step=tf.Variable(0,dtype=tf.int64,name='step');self.state=self.parameters+self.momentum+[self.step]
self.state_names=[v.name.split(':')[0] for v in self.state]
self.infer=tf.function(self._infer,input_signature=features)
self.train=tf.function(self._train,input_signature=features+[
tf.TensorSpec([1,3,S,S],tf.float32,name='y'),
tf.TensorSpec([1,2],tf.float32,name='presence'),
tf.TensorSpec([1],tf.float32,name='abstention'),
tf.TensorSpec([4],tf.float32,name='supervision'),
tf.TensorSpec([],tf.float32,name='learning_rate')])
def _infer(self,*f):
return {n:(v@self.pw+self.pb if n=='presence_logits' else v*self.scale[i]+self.bias[i]) for i,(n,v) in enumerate(zip(names,f))}
def _train(self,*args):
y,presence,abstention,supervision,rate=args[-5:]
with tf.GradientTape() as tape:
outputs=self._infer(*args[:-5]);valid=y[:,2:3]
spatial=[]
for name,target in [('segmentation_logits',y[:,0:1]),('point_logits',y[:,1:2])]:
loss=tf.nn.sigmoid_cross_entropy_with_logits(labels=target,logits=outputs[name])
spatial.append(tf.reduce_sum(loss*valid)/tf.maximum(tf.reduce_sum(valid),1.))
point,seg=spatial[1],spatial[0]
abst=tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(labels=abstention,logits=outputs['abstention_logits']))
pres=tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(labels=presence,logits=outputs['presence_logits']))
loss=tf.reduce_sum(tf.stack([seg,point,abst,pres])*supervision*tf.constant([1.,.5,.25,.25]))
gradients=tape.gradient(loss,self.parameters);gradients,_=tf.clip_by_global_norm(gradients,10.)
for v,m,g in zip(self.parameters,self.momentum,gradients):m.assign(.9*m+g);v.assign_sub(rate*m)
self.step.assign_add(1);return {'loss':loss,'step':self.step.read_value()}
@tf.function(input_signature=[tf.TensorSpec([],tf.string,name='checkpoint_path')])
def save(self,checkpoint_path):
tf.raw_ops.Save(filename=checkpoint_path,tensor_names=self.state_names,data=[v.read_value() for v in self.state]);return {'step':self.step.read_value()}
@tf.function(input_signature=[tf.TensorSpec([],tf.string,name='checkpoint_path')])
def restore(self,checkpoint_path):
for n,v in zip(self.state_names,self.state):v.assign(tf.raw_ops.Restore(file_pattern=checkpoint_path,tensor_name=n,dt=v.dtype))
return {'step':self.step.read_value()}
module=Head();signatures=['infer','train','save','restore'];saved=out/'head_saved_model'
tf.saved_model.save(module,str(saved),signatures={n:getattr(module,n).get_concrete_function() for n in signatures})
converter=tf.lite.TFLiteConverter.from_saved_model(str(saved));converter.target_spec.supported_ops=[tf.lite.OpsSet.TFLITE_BUILTINS,tf.lite.OpsSet.SELECT_TF_OPS];converter.experimental_enable_resource_variables=True
head=out/'head.tflite';head.write_bytes(converter.convert());base=root/'learning-backbones'/key/'model.tflite'
dynamic=out/'dynamic_backbone.tflite';preprocessing=dynamic_resize_backbone(base,dynamic)
artifact=out/'model.tflite';merge(dynamic,head,artifact,{'features_'+n:i for i,n in enumerate(names)},image_size=S)
def load():
i=tf.lite.Interpreter(model_path=str(artifact),num_threads=6,experimental_op_resolver_type=tf.lite.experimental.OpResolverType.BUILTIN_WITHOUT_DEFAULT_DELEGATES)
return i,{n:i.get_signature_runner(n) for n in signatures}
i,s=load();x=np.random.default_rng(94).normal(size=(1,3,S,S)).astype(np.float32);before=s['infer'](x=x)
reference=tf.lite.Interpreter(model_path=str(base),num_threads=6,experimental_op_resolver_type=tf.lite.experimental.OpResolverType.BUILTIN_WITHOUT_DEFAULT_DELEGATES);reference.allocate_tensors();reference.set_tensor(reference.get_input_details()[0]['index'],x);reference.invoke()
for name,d in zip(names,reference.get_output_details()):np.testing.assert_allclose(before[name],reference.get_tensor(d['index']),atol=2e-3,rtol=2e-3)
y=np.zeros((1,3,S,S),np.float32);y[:,2]=1;y[:,0,150:300,140:280]=1;y[:,1,210:230,200:220]=1
targets={'y':y,'presence':np.array([[1,0]],np.float32),'abstention':np.array([0],np.float32),'supervision':np.ones(4,np.float32)}
losses=[float(s['train'](x=x,**targets,learning_rate=np.float32(.01))['loss']) for _ in range(6)]
assert np.isfinite(losses).all() and losses[-1]<losses[0],losses
after=s['infer'](x=x);assert all(not np.array_equal(before[n],after[n]) for n in names)
checkpoint=np.array(str(out/'validation_checkpoint'),dtype=np.bytes_);s['save'](checkpoint_path=checkpoint)
j,t=load();t['restore'](checkpoint_path=checkpoint)
for n,v in after.items():np.testing.assert_array_equal(v,t['infer'](x=x)[n])
shapes=[(193,317),(321,257)]
for h,w in shapes:
xx=np.random.default_rng(4).normal(size=(1,3,h,w)).astype(np.float32)
for runner in (s,t):runner['train'](x=xx,**targets,learning_rate=np.float32(.001))
a=s['infer'](x=xx);b=t['infer'](x=xx)
for n in names:np.testing.assert_array_equal(a[n],b[n])
# Missing labels must not update an otherwise fresh adapter.
k,u=load();empty={**targets,'supervision':np.zeros(4,np.float32)};initial=u['infer'](x=x)
assert float(u['train'](x=x,**empty,learning_rate=np.float32(.01))['loss'])==0
for n,v in initial.items():np.testing.assert_array_equal(v,u['infer'](x=x)[n])
report={'key':key,'validation_passed':True,'training_scope':'four_task_output_adaptation_only','backbone_frozen':True,'limitation':'Frozen DINOv3 and DPT; only task-output calibration/mixing adapts. Pilot, private research use.','initial_detector_parity':True,'parameters_changed':True,'all_four_tasks_changed':True,'missing_supervision_no_update':True,'initial_loss':losses[0],'final_loss':losses[-1],'mixed_image_shapes':shapes,'restart_logits_exact':True,'optimizer_resume_exact':True,'variable_image_preprocessing':preprocessing,'android_device_tested':False,'accuracy_evaluated':False,'sha256':hashlib.sha256(artifact.read_bytes()).hexdigest(),'bytes':artifact.stat().st_size,'signatures':i.get_signature_list()}
(out/'training_validation.json').write_text(json.dumps(report,indent=2)+'\n')
training={'trainSignature':'train','inferSignature':'infer','saveSignature':'save','restoreSignature':'restore','imageInput':'x','targetInput':'y','lossOutput':'loss','checkpointInput':'checkpoint_path','learningRateInput':'learning_rate','inferOutputs':names,'targetEncoding':'segmentation_point_valid_mask_nchw','targetShape':[1,3,S,S],'auxiliaryTargets':{'presence':{'shape':[1,2],'labels':['flame_visible','smoke_visible']},'abstention':{'shape':[1]},'supervision':{'shape':[4],'order':['segmentation','point','abstention','presence'],'meaning':'1=reviewed label exists, 0=missing; not a negative label'}},'scope':report['training_scope']}
config={'task':'multitask','adapter':'fireviewer_dinov3_multitask','inputWidth':S,'inputHeight':S,'inputLayout':'NCHW','inputType':'FLOAT32','mean':0.,'std':1.,'channelMean':[123.675,116.28,103.53],'channelStd':[58.395,57.12,57.375],'resizeMode':'stretch','isRgb':True,'labels':['flame_visible','smoke_visible'],'outputLayout':'NCHW','scoreActivation':'sigmoid','dynamicMinSize':128,'dynamicMaxSize':1280,'dynamicStride':1,'training':training}
(out/'android_model_config.json').write_text(json.dumps(config,indent=2)+'\n');(out/'source.json').write_text(json.dumps({'key':key,'kind':'multitask','labels':config['labels'],'backbone':str(base),'private_pilot_only':True},indent=2)+'\n');print(json.dumps(report),flush=True)