unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified Download android/TrainingRuntimeCheck.java from fireviewer/litert-models: direct link, hf CLI and curl.
- Browser
- Download file 3.25 kB
-
https://huggingface.co/fireviewer/litert-models/resolve/main/android/TrainingRuntimeCheck.java
- Command line
-
hf download hf://fireviewer/litert-models/android/TrainingRuntimeCheck.java
-
curl -L -o TrainingRuntimeCheck.java https://huggingface.co/fireviewer/litert-models/resolve/main/android/TrainingRuntimeCheck.java
3.25 kB
| import java.io.File; | |
| import java.util.Arrays; | |
| import java.util.HashMap; | |
| import java.util.Map; | |
| import java.nio.ByteBuffer; | |
| import java.nio.ByteOrder; | |
| import org.tensorflow.lite.Interpreter; | |
| /** Runs under Android app_process with the matched TensorFlow Lite 2.16.1 AARs. */ | |
| public final class TrainingRuntimeCheck { | |
| static float[][][][] image(int h,int w) { | |
| float[][][][] x=new float[1][3][h][w]; | |
| for(int c=0;c<3;c++)for(int y=0;y<h;y++)for(int k=0;k<w;k++)x[0][c][y][k]=(float)Math.sin(c+y*.07+k*.11); | |
| return x; | |
| } | |
| static Map<String,Object> inputs(Object... values) { | |
| Map<String,Object> out=new HashMap<>();for(int n=0;n<values.length;n+=2)out.put((String)values[n],values[n+1]);return out; | |
| } | |
| static float[] infer(Interpreter i,float[][][][] x,int count) { | |
| float[][] scores=new float[1][count];i.runSignature(inputs("x",x),inputs("scores",scores),"infer");return scores[0]; | |
| } | |
| static float train(Interpreter i,float[][][][] x,float[][] y) { | |
| ByteBuffer rate=ByteBuffer.allocateDirect(4).order(ByteOrder.nativeOrder());rate.putFloat(0,.0001f); | |
| ByteBuffer loss=ByteBuffer.allocateDirect(4).order(ByteOrder.nativeOrder()); | |
| i.runSignature(inputs("x",x,"y",y,"learning_rate",rate),inputs("loss",loss),"train");return loss.getFloat(0); | |
| } | |
| static void check(boolean okay,String message){if(!okay)throw new AssertionError(message);} | |
| public static void main(String[] args) { | |
| File model=new File(args[0]);String checkpoint=args[1];int classes=Integer.parseInt(args[2]); | |
| Interpreter.Options options=new Interpreter.Options().setNumThreads(2).setUseXNNPACK(false); | |
| float[][][][] x=image(224,224);float[][] target=new float[1][classes];target[0][7]=1; | |
| float[] trained,continued;float first=0,last=0; | |
| System.out.println("ANDROID_CHECK opening interpreter"); | |
| try(Interpreter i=new Interpreter(model,options)) { | |
| System.out.println("ANDROID_CHECK interpreter open"); | |
| float[] original=infer(i,x,classes); | |
| System.out.println("ANDROID_CHECK inference complete"); | |
| for(int n=0;n<6;n++){last=train(i,n%2==0?x:image(192,320),target);if(n==0)first=last;check(Float.isFinite(last),"loss nonfinite");System.out.println("ANDROID_CHECK step="+n+" loss="+last);} | |
| trained=infer(i,x,classes);check(!Arrays.equals(original,trained),"weights did not alter output"); | |
| i.runSignature(inputs("checkpoint_path",checkpoint),new HashMap<>(),"save"); | |
| check(new File(checkpoint).length()>0,"empty checkpoint"); | |
| train(i,x,target);continued=infer(i,x,classes); | |
| } | |
| try(Interpreter restored=new Interpreter(model,options)) { | |
| restored.runSignature(inputs("checkpoint_path",checkpoint),new HashMap<>(),"restore"); | |
| check(Arrays.equals(trained,infer(restored,x,classes)),"restore differs"); | |
| train(restored,x,target);check(Arrays.equals(continued,infer(restored,x,classes)),"optimizer resume differs"); | |
| } | |
| System.out.println("ANDROID_TRAINING_PASS {\"model\":\""+model.getName()+"\",\"mixed_shapes\":[[224,224],[192,320]],\"first_loss\":"+first+",\"last_loss\":"+last+",\"restart_exact\":true,\"optimizer_resume_exact\":true,\"runtime\":\""+org.tensorflow.lite.TensorFlowLite.runtimeVersion()+"\",\"device\":\"Android emulator x86_64\"}"); | |
| } | |
| } | |