1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
type t = { weights : float array; bias : float }
type example = float array * float
let make ~(inputs : int) ~(seed : int) : t =
let st = Lehmer.make seed in
{ weights = Array.init inputs (fun _ -> Lehmer.float st 0.2 -. 0.1); bias = 0. }
let sum (n : t) (x : float array) : float =
if Array.length x <> Array.length n.weights then invalid_arg "Neuron: wrong number of inputs";
let s = ref n.bias in
Array.iteri (fun i v -> s := !s +. (n.weights.(i) *. v)) x;
!s
let answer (n : t) (x : float array) : float = if sum n x > 0. then 1. else 0.
let learn ?(rate = 0.1) (n : t) ((x, target) : example) : t =
let wrong = target -. answer n x in
if wrong = 0. then n
else
{ weights = Array.mapi (fun i w -> w +. (rate *. wrong *. x.(i))) n.weights; bias = n.bias +. (rate *. wrong) }
let epoch ?rate (n : t) (examples : example list) : t = List.fold_left (fun n e -> learn ?rate n e) n examples
let mistakes (n : t) (examples : example list) : int =
List.length (List.filter (fun ((x, target) : example) -> answer n x <> target) examples)
let train ?(epochs = 100) ?rate (n : t) (examples : example list) : t =
let rec go n left = if left = 0 || mistakes n examples = 0 then n else go (epoch ?rate n examples) (left - 1) in
go n epochs
let learns ?epochs ?(seed = 1) (examples : example list) : float =
match examples with
| [] -> 1.
| (x, _) :: _ ->
let n = train ?epochs (make ~inputs:(Array.length x) ~seed) examples in
let got = List.length examples - mistakes n examples in
float_of_int got /. float_of_int (List.length examples)
let problem (f : float -> float -> float) : example list =
List.map (fun (a, b) -> ([| a; b |], f a b)) [ (0., 0.); (0., 1.); (1., 0.); (1., 1.) ]
let and_ : example list = problem (fun a b -> if a = 1. && b = 1. then 1. else 0.)
let or_ : example list = problem (fun a b -> if a = 1. || b = 1. then 1. else 0.)
let xor : example list = problem (fun a b -> if a <> b then 1. else 0.)