# frozen_string_literal: true
require 'rubygems'
gem 'ruby-fann', '= 2.0.2.oneruby1'
require 'ruby-fann'
require 'json'
require 'tmpdir'

module XorExperiment
  INPUTS=[[0.0,0.0],[0.0,1.0],[1.0,0.0],[1.0,1.0]].freeze
  TARGETS=[[0.0],[1.0],[1.0],[0.0]].freeze
  def self.predict(network, values)
    unless values.is_a?(Array) && values.length==2 && values.all? { |v| v.is_a?(Numeric) && !v.is_a?(Complex) && v.finite? && v.between?(0,1) }
      raise ArgumentError, 'two finite numbers in [0, 1] required'
    end
    network.run(values.map(&:to_f)).first
  end
  def self.train
    data=RubyFann::TrainData.new(inputs: INPUTS, desired_outputs: TARGETS)
    # FANN initializes randomly. Bound the work and reject an unsuccessful fit.
    5.times do
      network=RubyFann::Standard.new(num_inputs: 2, hidden_neurons: [4], num_outputs: 1)
      10_000.times do |epoch|
        network.train_epoch(data)
        next unless epoch % 25 == 0
        predictions=INPUTS.map { |row| predict(network,row) }
        mse=predictions.zip(TARGETS.flatten).sum { |p,y| (p-y)**2 } / INPUTS.length
        return network if mse < 0.0025 && predictions.map { |p| p >= 0.5 ? 1 : 0 } == [0,1,1,0]
      end
    end
    raise 'no acceptable XOR fit within five bounded attempts'
  end
end
if $PROGRAM_NAME==__FILE__
  net=XorExperiment.train
  predictions=XorExperiment::INPUTS.map { |row| XorExperiment.predict(net,row) }
  Dir.mktmpdir do |dir|
    path=File.join(dir,'xor.net');net.save(path)
    loaded=RubyFann::Standard.new(filename: path)
    reloaded=XorExperiment::INPUTS.map { |row| XorExperiment.predict(loaded,row) }
    puts JSON.generate(classes: predictions.map { |p| p>=0.5 ? 1 : 0 },
                       mse: predictions.zip([0,1,1,0]).sum { |p,y| (p-y)**2 }/4,
                       reload_max_error: predictions.zip(reloaded).map { |a,b| (a-b).abs }.max)
  end
end
