litert-models / conversion /build_trainable_detector.py
unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified
Raw History Blame Contribute Delete
12.1 kB
"""Frozen detector plus trainable candidate classification/box adaptation on Android.
The adaptation cannot generate proposals absent from the frozen detector. It is
not full detector fine-tuning. Targets are reviewed xyxy/class/mask annotations.
"""
import os
os.environ.update(TF_CPP_MIN_LOG_LEVEL='2',TF_NUM_INTRAOP_THREADS='6',TF_NUM_INTEROP_THREADS='2')
import sys,json,hashlib
from pathlib import Path
import tensorflow as tf,numpy as np
from merge_training_graph import merge,dynamic_resize_backbone,freeze_auxiliary_inputs
root=Path('/workspace/vds-litert-batch2');key=sys.argv[1];out=root/'learning'/key;out.mkdir(parents=True,exist_ok=True)
source=json.loads((out/'source.json').read_text());C=source['class_count'];kind=source['kind']
class_indices=source.get('class_indices',list(range(C)))
specs=([tf.TensorSpec([1,C+4,None],tf.float32,name='features')] if kind=='yolo' else [tf.TensorSpec([1,None,4],tf.float32,name='features_boxes'),tf.TensorSpec([1,None,C],tf.float32,name='features_logits')])
if kind=='rtmdet':specs=[tf.TensorSpec([1,None,None,C+4],tf.float32,name=f'features_{stride}') for stride in (8,16,32)]
specs+=[tf.TensorSpec([1,3,None,None],tf.float32,name='image')]
class Adaptation(tf.Module):
def __init__(self):
super().__init__()
self.cw=tf.Variable(tf.eye(C),name='class_weight');self.cb=tf.Variable(tf.zeros([C]),name='class_bias')
self.bw=tf.Variable(tf.eye(4),name='box_weight');self.bb=tf.Variable(tf.zeros([4]),name='box_bias')
self.parameters=[self.cw,self.cb,self.bw,self.bb]
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.names=['class_weight','class_bias','box_weight','box_bias','momentum_class_weight','momentum_class_bias','momentum_box_weight','momentum_box_bias','step']
self.infer=tf.function(self._infer,input_signature=specs)
self.train=tf.function(self._train,input_signature=specs+[tf.TensorSpec([1,100,6],tf.float32,name='y'),tf.TensorSpec([],tf.float32,name='learning_rate')])
def predict(self,args):
image=args[-1];shape=tf.cast(tf.shape(image),tf.float32);scale=tf.stack([shape[3],shape[2],shape[3],shape[2]])
if kind=='yolo':
values=tf.transpose(args[0],[0,2,1]);boxes=values[:,:,:4]/scale
p=tf.clip_by_value(values[:,:,4:],1e-6,1-1e-6);logits=tf.math.log(p)-tf.math.log1p(-p)
elif kind=='rtmdet':
all_boxes=[];all_logits=[]
for feature,stride in zip(args[:3],(8,16,32)):
fs=tf.shape(feature);grid_x,grid_y=tf.meshgrid(tf.range(fs[2]),tf.range(fs[1]));grid=tf.cast(tf.stack([grid_x,grid_y],-1),tf.float32)*stride
d=feature[:,:,:,C:]*stride;lo=grid[None]-d[:,:,:,:2];hi=grid[None]+d[:,:,:,2:]
boxes=tf.concat([(lo+hi)/2,hi-lo],-1)/scale
all_boxes.append(tf.reshape(boxes,[1,-1,4]));all_logits.append(tf.reshape(feature[:,:,:,:C],[1,-1,C]))
boxes=tf.concat(all_boxes,1);logits=tf.concat(all_logits,1)
else:boxes,logits=args[:2]
return boxes@self.bw+self.bb,logits@self.cw+self.cb,scale
def _infer(self,*args):
boxes,logits,scale=self.predict(args)
if kind=='yolo':return {'detections':tf.transpose(tf.concat([boxes*scale,tf.sigmoid(logits)],-1),[0,2,1])}
probabilities=tf.nn.softmax(logits,axis=-1) if source.get('class_activation')=='softmax' else tf.sigmoid(logits)
probabilities=tf.gather(probabilities,class_indices,axis=-1)
return {'boxes':boxes,'logits':logits,'detections':tf.transpose(tf.concat([boxes*scale,probabilities],-1),[0,2,1])}
def _train(self,*args):
y,lr=args[-2:];targets=y[0];valid=targets[:,5];gt=targets[:,:4];labels=tf.gather(tf.constant(class_indices,tf.int32),tf.cast(targets[:,4],tf.int32))
with tf.GradientTape() as tape:
boxes,logits,_=self.predict(args[:-2]);b=boxes[0];scores=logits[0]
xyxy=tf.concat([b[:,:2]-b[:,2:]/2,b[:,:2]+b[:,2:]/2],-1)
lo=tf.maximum(xyxy[:,None,:2],gt[None,:,:2]);hi=tf.minimum(xyxy[:,None,2:],gt[None,:,2:])
inter=tf.reduce_prod(tf.maximum(hi-lo,0),-1)
area=tf.reduce_prod(tf.maximum(xyxy[:,2:]-xyxy[:,:2],0),-1)
ga=tf.reduce_prod(tf.maximum(gt[:,2:]-gt[:,:2],0),-1)
iou=inter/tf.maximum(area[:,None]+ga[None,:]-inter,1e-7)*valid[None,:]
# Assignment is discrete; gradients update the prediction, never the matching.
best_candidate=tf.argmax(iou,axis=0,output_type=tf.int32)
forced_pairs=tf.transpose(tf.one_hot(best_candidate,tf.shape(b)[0])*valid[:,None])
forced=tf.reduce_max(forced_pairs,axis=1)
# Assign a forced candidate to the annotation that selected it,
# including zero-IoU cases. Collisions keep the larger-IoU target.
best_target=tf.argmax(iou+forced_pairs*(1+iou),axis=1,output_type=tf.int32)
positive=tf.stop_gradient(tf.maximum(tf.cast(tf.reduce_max(iou,axis=1)>.5,tf.float32),forced))
assigned=tf.gather(gt,best_target);assigned_class=tf.gather(labels,best_target)
onehot=tf.one_hot(assigned_class,C)*positive[:,None]
if source.get('class_activation')=='softmax':
onehot+=tf.one_hot(source['background_index'],C)[None,:]*(1-positive[:,None])
ce=tf.nn.softmax_cross_entropy_with_logits(labels=onehot,logits=scores)
else:ce=tf.reduce_mean(tf.nn.sigmoid_cross_entropy_with_logits(labels=onehot,logits=scores),axis=-1)
cls=tf.reduce_sum(ce*positive)/tf.maximum(tf.reduce_sum(positive),1)+.05*tf.reduce_sum(ce*(1-positive))/tf.maximum(tf.reduce_sum(1-positive),1)
target_box=tf.concat([(assigned[:,:2]+assigned[:,2:])/2,assigned[:,2:]-assigned[:,:2]],-1)
geometry=tf.reduce_sum(tf.reduce_mean(tf.square(b-target_box),axis=-1)*positive)/tf.maximum(tf.reduce_sum(positive),1)
loss=cls+5*geometry
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(lr*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.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.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()}
@tf.function(input_signature=[])
def training_state(self):return {'head_kernel':self.cw.read_value(),'step':self.step.read_value()}
module=Adaptation();names=['infer','train','save','restore','training_state'];saved=out/'head_saved_model'
tf.saved_model.save(module,str(saved),signatures={n:getattr(module,n).get_concrete_function() for n in names})
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());artifact=out/'model.tflite';backbone=Path(source['backbone'])
if source.get('freeze_auxiliary_inputs'):
frozen=out/'frozen_inputs.tflite';freeze_auxiliary_inputs(backbone,frozen,source['freeze_auxiliary_inputs']);backbone=frozen
preprocessing={'mode':'native_dynamic','stride':32}
if not source.get('native_dynamic',True):
wrapped=out/'dynamic_backbone.tflite';preprocessing=dynamic_resize_backbone(backbone,wrapped);backbone=wrapped
merge(backbone,head,artifact,source['bindings'],image_size=source.get('app_size',320))
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 names}
i,s=load();h,w=source.get('test_size',[320,320]);x=np.random.default_rng(34).uniform(0,1,(1,3,h,w)).astype(np.float32)
before=s['infer'](x=x)
# Initial adaptation must preserve detector predictions.
base=tf.lite.Interpreter(model_path=str(backbone),num_threads=6,experimental_op_resolver_type=tf.lite.experimental.OpResolverType.BUILTIN_WITHOUT_DEFAULT_DELEGATES)
d=base.get_input_details()[0];base.resize_tensor_input(d['index'],x.shape,strict=True);base.allocate_tensors();base.set_tensor(d['index'],x);base.invoke()
raw=[base.get_tensor(d['index']) for d in base.get_output_details()]
if kind=='yolo':
np.testing.assert_allclose(before['detections'],raw[source['bindings']['features']],atol=1e-3,rtol=1e-3)
pred=before['detections'][0].T;index=np.argmax(pred[:,4:].max(-1));box=pred[index,:4]/np.array([w,h,w,h]);label=int(np.argmax(pred[index,4:]))
elif kind=='rtmdet':
bb=[];ss=[]
for stride in (8,16,32):
feature=raw[source['bindings'][f'features_{stride}']];yy,xx=np.mgrid[:feature.shape[1],:feature.shape[2]];grid=np.stack([xx,yy],-1)*stride
d=feature[:,:,:,C:]*stride;lo=grid[None]-d[:,:,:,:2];hi=grid[None]+d[:,:,:,2:]
bb.append((np.concatenate([(lo+hi)/2,hi-lo],-1)/np.array([w,h,w,h])).reshape(1,-1,4));ss.append(feature[:,:,:,:C].reshape(1,-1,C))
np.testing.assert_allclose(before['boxes'],np.concatenate(bb,1),atol=1e-4,rtol=1e-4);np.testing.assert_allclose(before['logits'],np.concatenate(ss,1),atol=1e-5,rtol=1e-5)
index=np.argmax(before['logits'][0].max(-1));box=before['boxes'][0,index];label=int(np.argmax(before['logits'][0,index]))
else:
np.testing.assert_allclose(before['boxes'],raw[source['bindings']['features_boxes']],atol=1e-5,rtol=1e-5)
np.testing.assert_allclose(before['logits'],raw[source['bindings']['features_logits']],atol=1e-5,rtol=1e-5)
candidates=before['logits'][0][:,class_indices];index=np.argmax(candidates.max(-1));box=before['boxes'][0,index];label=int(np.argmax(candidates[index]))
y=np.zeros((1,100,6),np.float32);y[0,0,:4]=np.clip(np.r_[box[:2]-box[2:]/2,box[:2]+box[2:]/2],0,1);y[0,0,4:]=[label,1]
losses=[]
for _ in range(8):losses.append(float(s['train'](x=x,y=y,learning_rate=np.float32(.005))['loss']))
assert np.isfinite(losses).all() and losses[-1]<losses[0],losses
after=s['infer'](x=x);assert any(not np.array_equal(before[n],after[n]) for n in before)
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=[(h,w),(h,w+64),(h+64,w)]
for hh,ww in shapes:
xx=np.random.default_rng(77).uniform(0,1,(1,3,hh,ww)).astype(np.float32)
for runner in (s,t):runner['train'](x=xx,y=y,learning_rate=np.float32(.001))
a=s['infer'](x=xx);b=t['infer'](x=xx)
for n in a:np.testing.assert_array_equal(a[n],b[n]);assert np.isfinite(a[n]).all()
report={'key':key,'validation_passed':True,'training_scope':'candidate_classification_and_box_adaptation_only','backbone_frozen':True,'limitation':'Does not create proposals absent from frozen detector','initial_detector_parity':True,'parameters_changed':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');print(json.dumps(report),flush=True)