
(declare (usual-integrations))

(define *svm-verbose* #f)

;; Default values for training parameters
(define *svm-max-alpha* 0.05)
(define *svm-tolerance* 0.001)
(define *svm-eps* 0.001)

;;; Controls drawing
(define *svm-draw-interval-while-training* 500)

(define *svm* #f)

;;; Constructs the svm-x and set-svm-x! functions for each of the
;;; components of the SVM data structure.  The disadvantage is that
;;; these generated functions will not get compiled.
(define-structure (svm)
  n-points n-supports dim kernel-spec alphas weights self-dots
  threshold data-points targets class0 class1)

;;; Used during training to avoid repeated work.
(define *svm-error-cache* #f)
(define *svm-precomputed-self-dot-product* #f)

;;; Pseudonyms for accessing the default *svm*
(define-integrable (svm-get-n-points) (svm-n-points *svm*)) 
(define-integrable (svm-get-n-supports) (svm-n-supports *svm*)) 
(define-integrable (svm-get-dim) (svm-dim *svm*)) 
(define-integrable (svm-get-kernel-spec) (svm-kernel-spec *svm*)) 
(define-integrable (svm-get-threshold) (svm-threshold *svm*)) 
(define-integrable (svm-get-class0) (svm-class0 *svm*)) 
(define-integrable (svm-get-class1) (svm-class1 *svm*)) 

;;; Some nicer ways of getting at the components of vector parameters
(define-integrable (svm-get-target i) (vector-ref (svm-targets *svm*) i))
(define-integrable (svm-set-target i v) (vector-set! (svm-targets *svm*) i v))

(define-integrable (svm-get-alpha i) (vector-ref (svm-alphas *svm*) i))
(define-integrable (svm-set-alpha i a) (vector-set! (svm-alphas *svm*) i a))

(define-integrable (svm-get-data-point i) (vector-ref (svm-data-points *svm*) i))
(define-integrable (svm-set-data-point i p) (vector-set! (svm-data-points *svm*) i p))

(define-integrable (svm-get-weight i) (vector-ref (svm-weights *svm*) i))
(define-integrable (svm-set-weight i w) (vector-set! (svm-weights *svm*) i w))

(define-integrable (svm-get-self-dot i) (vector-ref (svm-self-dots *svm*) i))
(define-integrable (svm-set-self-dot i d) (vector-set! (svm-self-dots *svm*) i d))

(define-integrable (svm-get-error-cache i) (vector-ref *svm-error-cache* i))
(define-integrable (svm-set-error-cache i e) (vector-set! *svm-error-cache* i e))

(define (svm-get-e i alph y)
  (if (and (> alph  0) (< alph *svm-max-alpha*))
      (svm-get-error-cache i)
      (- (svm-learned-func i) y)))

(define (examine-example i1)
  (let* ((alph1 (svm-get-alpha i1))
	 (y1 (svm-get-target i1))
	 (e1 (svm-get-e i1 alph1 y1))
	 (r1 (* y1 e1)))
    
    (define (heuristic1 k i2 tmax)
      (if (= k (svm-get-n-points))
	  (if (and (>= i2 0) (take-step i1 i2))
	      1
	      #f)
	  (let ((alph_k (svm-get-alpha k)))
	    (cond ((and (> alph_k 0) (< alph_k *svm-max-alpha*))
		   (if *svm-verbose* (display* 'heuristc1 " " k " " i2 " " tmax))
		   (let* ((e2 (svm-get-error-cache k))
			  (temp (abs (- e1 e2))))
		     (if (> temp tmax)
			 (heuristic1 (+ k 1) k temp)
			 (heuristic1 (+ k 1) i2 tmax))))
		  (else 
		   (heuristic1 (+ k 1) i2 tmax))))))

    (define (heuristic2 k end)
      (if (= k end)
	  #f
	  (let* ((i2 (modulo k (svm-get-n-points)))
		 (alph-i2 (svm-get-alpha i2)))
	    (cond ((and (> alph-i2 0) (< alph-i2 *svm-max-alpha*))
		   (if *svm-verbose* (display* 'heuristc2 " " k " " end))
		   (if (take-step i1 i2)
		       1
		       (heuristic2 (+ k 1) end)))
		  (else
		   (heuristic2 (+ k 1) end))))))

    (define (heuristic3 k end)
      (if (= k end)
	  #f
	  (cond ((take-step i1 (modulo k (svm-get-n-points)))
		 (if *svm-verbose* (display* 'heuristc3 " " k " " end))
		 1)
		(else
		 (heuristic3 (+ k 1) end)))))

    (if *svm-verbose*
	(display* "alpha=" (svm-alphas *svm*) "b=" (svm-get-threshold) "ec=" *svm-error-cache*))

    (cond ((or (and (< r1 (- *svm-tolerance*)) (< alph1 *svm-max-alpha*))
	       (and (> r1 *svm-tolerance*) (> alph1 0)))
	   ;; Try i2 by three heuristics, if successful, then immediately return 1
	   (or
	    (heuristic1 0 -1 0)
	    (let ((k0 (random (svm-get-n-points))))
	      (heuristic2 k0 (+ (svm-get-n-points) k0)))
	    (let ((k0 (random (svm-get-n-points))))
	      (heuristic3 k0 (+ (svm-get-n-points) k0)))
	    ;; nothing worked.
	    0
	    ))
	  (else 0))
    ))

(define (take-step i1 i2) 
  (cond 
   ((= i1 i2) 
    ;; return #f
    #f)
   (else
    (let* ((alph1 (svm-get-alpha i1))
	   (y1 (svm-get-target i1))
	   (e1 (svm-get-e i1 alph1 y1))
	   (alph2 (svm-get-alpha i2))
	   (y2 (svm-get-target i2))
	   (e2 (svm-get-e i2 alph2 y2))
	   (s (* y1 y2))
	   (gamma (if (= y1 y2) (+ alph1 alph2) (- alph1 alph2)))
	   (L (if (= y1 y2)
		  (if (> gamma *svm-max-alpha*) (- gamma *svm-max-alpha*) 0)
		  (if (> gamma 0) 0 (- gamma))))
	   (H (if (= y1 y2)
		  (if (> gamma *svm-max-alpha*) *svm-max-alpha* gamma)
		  (if (> gamma 0) (- *svm-max-alpha* gamma) *svm-max-alpha*))))
      (if *svm-verbose* (display* "e1=" e1 " e2=" e2 " s=" s " L=" L  " H=" H))
      (cond ((= L H)
	     (if *svm-verbose* (display* "L=H"))
	     ;; return #f
	     #f)
	    (else
	     (let* ((k11 (svm-kernel-func i1 i1))
		    (k12 (svm-kernel-func i1 i2))
		    (k22 (svm-kernel-func i2 i2))
		    (eta (- (* 2 k12) k11 k22))
		    (a2
		     (if (< eta 0)
			 (max L (min H (+ alph2 (/ (* y2 (- e2 e1)) eta))))
			 (let* ((c1 (/ eta 2))
				(c2 (- (* y2 (- e1 e2)) (* eta alph2)))
				(Lobj (+ (* c1 L L) (* c2 L)))
				(Hobj (+ (* c1 H H) (* c2 H))))
			   (cond ((> Lobj (+ Hobj *svm-eps*)) L)
				 ((< Lobj (- Hobj *svm-eps*)) H)
				 (else alph2))))))
	       (cond ((< (abs (- a2 alph2)) (* *svm-eps* (+ a2 alph2 *svm-eps*)))
		      (if *svm-verbose*
			  (display* "(< (abs (- a2 alph2)) (* *svm-eps* (+ a2 alph2 *svm-eps*)))"))
		      ;; return #f
		      #f)
		     (else
		      (let* ((a1t (- alph1 (* s (- a2 alph2))))
			     (a1 (max 0 (min *svm-max-alpha* a1t))))
			(if (< a1t 0) 
			    (set! a2 (+ a2 (* s a1)))
			    (if (> a1t *svm-max-alpha*)
				(set! a2 (+ a2 (* s (- a1 *svm-max-alpha*))))))
			(let* ((t1 (* y1 (- a1 alph1)))
			       (t2 (* y2 (- a2 alph2)))
			       (delta-b (update-b e1 e2 a1 a2 t1 t2 k11 k12 k22)))
			  (update-weights i1 i2 t1 t2 delta-b)
			  (svm-set-alpha i1 a1)
			  (svm-set-alpha i2 a2)
			  ;; return #t
			  #t))))
	       ))))))
  )

(define (update-weights i1 i2 t1 t2 delta-b)
  (if (eq? (first (svm-get-kernel-spec)) 'linear)
      (do ((i 0 (1+ i))
	   (p1 (svm-get-data-point i1) (cdr p1))
	   (p2 (svm-get-data-point i2) (cdr p2)))
	  ((= i (svm-get-dim)))
	(svm-set-weight i 
			(+ (svm-get-weight i)
			   (+ (* t1 (first p1)) (* t2 (first p2)))))))
  (do ((i 0 (1+ i)))
      ((= i (svm-get-n-points))
       (svm-set-error-cache i1 0)
       (svm-set-error-cache i2 0)
       )
    (let ((alph-i (svm-get-alpha i)))
      (if (and (> alph-i 0) (< alph-i *svm-max-alpha*))
	  (svm-set-error-cache i 
			       (+ (svm-get-error-cache i)
				  (+ (* (svm-kernel-func i1 i) t1) 
				     (* (svm-kernel-func i2 i) t2)
				     (- delta-b))
				  ))))))

(define (update-b e1 e2 a1 a2 t1 t2 k11 k12 k22)
  (let ((bold (svm-get-threshold))
	(bnew
	 (if (and (> a1 0) (< a1 *svm-max-alpha*))
	     (+ (svm-get-threshold) e1 (* t1 k11) (* t2 k12))
	     (if (and (> a2 0) (< a2 *svm-max-alpha*))
		 (+ (svm-get-threshold) e2 (* t1 k12) (* t2 k22))
		 (let ((b1 (+ (svm-get-threshold) e1 (* t1 k11) (* t2 k12)))
		       (b2 (+ (svm-get-threshold) e2 (* t1 k12) (* t2 k22))))
		   (/ (+ b1 b2) 2.))))))
    (set-svm-threshold! *svm* bnew)
    ;; return delta-b
    (- bnew bold)))

(define (learned-func-linear k) 
  (do ((i 0 (1+ i))
       (pk (svm-get-data-point k) (cdr pk))
       (s 0 (+ s (* (svm-get-weight i) (first pk)))))
      ((= i (svm-get-dim)) (- s (svm-get-threshold))))
  )

(define (learned-func-nonlinear k) 
  (do ((i 0 (1+ i))
       (s 0 (if (> (svm-get-alpha i) 0)
		(+ s (* (svm-get-alpha i) (svm-get-target i)
			(svm-kernel-func i k)))
		s)))
      ((= i (svm-get-n-supports)) (- s (svm-get-threshold))))
  )

(define (svm-learned-func k)
  (let ((kernel-spec (svm-get-kernel-spec)))
    (cond ((eq? (first kernel-spec) 'linear) (learned-func-linear k))
	  ((memq (first kernel-spec) '(gauss poly)) (learned-func-nonlinear k))
	  (else (error "Unsupported kernel: " kernel-spec)))))

(define (dot-product i1 i2)
  (apply + (map * (svm-get-data-point i1) (svm-get-data-point i2))))

(define (rbf-kernel i1 i2 two-sigma-squared)
  ;; (p1 - p2).(p1 - p2) = p1.p1 -2 p1.p2 + p2.p2
  (exp (/ (- (+ (svm-get-self-dot i1)
		(* -2.0 (dot-product i1 i2))
		(svm-get-self-dot i2)))
	  ;; kernel spec is (rbf two-sigma-squared)
	  two-sigma-squared)))

(define (poly-kernel i1 i2 d)
  (expt (+ 1 (dot-product i1 i2)) d))

(define (svm-kernel-func i1 i2)
  (let ((kernel-spec (svm-get-kernel-spec)))
    (cond ((eq? (first kernel-spec) 'linear) 
	   (dot-product i1 i2))
	  ((eq? (first kernel-spec) 'gauss)
	   (rbf-kernel i1 i2 (second (svm-get-kernel-spec))))
	  ((eq? (first kernel-spec) 'poly)
	   (poly-kernel i1 i2 (second (svm-get-kernel-spec))))
	  (else (error "Unsupported kernel: " kernel-spec)))))

(define (svm-set-classes data)
  (let ((classes (map car (class-counts data))))
    (cond ((= (length classes) 2)
	   (cond ((member classes '((0 1) (0.0 1.0) (1 0) (1.0 0.0)))
		  (set-svm-class0! *svm* 0)
		  (set-svm-class1! *svm* 1))
		 ((member classes '((-1 1) (-1.0 1.0) (-1 0) (1.0 -1.0)))
		  (set-svm-class0! *svm* -1)
		  (set-svm-class1! *svm* 1))
		 (else
		  (set-svm-class0! *svm* (first classes))
		  (set-svm-class1! *svm* (second classes)))))
	  (else 
	   (error "Only two classes allowed, but we have" classes)))))

(define (svm-setup data kernel-spec)
  ;; Set up the SVM data structures
  (let ((n (length data)))
    (set! *svm* 
	  (make-svm
	   n				; n-points
	   n				; n-supports
	   (length (data-point-features (first data))) ; dim
	   kernel-spec			; kernel-spec
	   (make-vector (svm-n-points *svm*) 0.0) ; alphas
	   ;; weights
	   (if (eq? (first kernel-spec) 'linear)
	       (make-vector (svm-dim *svm*) 0.0)
	       #f)
	   ;; self dots
	   (cond ((not (eq? (first kernel-spec) 'linear))
		  (make-vector (+ 1 (svm-n-points *svm*)))
		  (do ((i 0 (+ i 1)))
		      ((= i (svm-n-points *svm*)))
		    (svm-set-self-dot i (dot-product i i))))
		 (else #f))
	   0				; threshold
	   (make-vector (+ 1 (svm-n-points *svm*))) ;data points
	   (make-vector (svm-n-points *svm*)) ; targets
	   #f				; class0
	   #f				; class1
	   ))
    (svm-set-classes data)
    (do ((i 0 (+ i 1))
	 (p data (cdr p)))
	((null? p))
      (svm-set-data-point i (data-point-features (car p)))
      (let ((class (data-point-class (car p))))
	(svm-set-target i (if (eq? class (svm-get-class0)) -1 1)))
      )
    ))

(define (svm-train-setup data kernel-func)
  (svm-setup data kernel-func)
  (set! *svm-error-cache* (make-vector (svm-get-n-points) 0.)))

(define (svm-train data kernel-spec)
  (svm-train-setup data kernel-spec)
  (do ((cycle 0 (+ 1 cycle))
       (num-changed 0)
       (examine-all #t))
      ((not (or (> num-changed 0) examine-all)))
    (set! num-changed 0)
    (if examine-all
	(do ((k 0 (+ 1 k)))
	    ((= k (svm-get-n-points)))
	  (set! num-changed (+ num-changed (examine-example k))))
	(do ((k 0 (+ 1 k)))
	    ((= k (svm-get-n-points)))
	  (let ((alph-k (svm-get-alpha k)))
	    (if (and (not (= alph-k 0)) (not (= alph-k *svm-max-alpha*)))
		(set! num-changed (+ num-changed (examine-example k)))))))
    (if examine-all
	(set! examine-all #f)
	(if (= num-changed 0)
	    (set! examine-all #t)))

    (do ((i 0 (+ i 1))
	 (bound-support 0)
	 (non-bound-support 0))
	((= i (svm-get-n-points))
	 (display* "Non-bound=" non-bound-support ", bound=" bound-support))
      (let ((alph-i (svm-get-alpha i))) 
	;; if alpha = 0, it's not a support at all...
	(if (> alph-i 0)
	    (if (< alph-i *svm-max-alpha*)
		(set! non-bound-support (+ 1 non-bound-support))
		(set! bound-support (+ 1 bound-support)))))
      )

      (if (and *draw* *svm-draw-interval-while-training*)
	  (cond ((zero? (remainder cycle *svm-draw-interval-while-training*))
		 (if (= cycle 0)
		     (draw-classifier svm-get-point data)
		     (draw-classifier svm-get-point data *window*)))))
    )
  
  (do ((i 0 (+ i 1))
       (errors 0))
      ((= i (svm-get-n-points))
       (display* "Got " errors " prediction errors.  Error rate = " 
		 (exact->inexact (/ errors (svm-get-n-points)))))
    (if (< (* (svm-learned-func i) (svm-get-target i)) 0) 
	(set! errors (+ 1 errors))))

  (if *draw* (draw-classifier svm-get-point data *window*))
  
  'done
  )

(define (svm-classify data-point)
  (let* ((out (make-svm-prediction (data-point-features data-point)))
	 (prediction (if (>= out 0) *svm-class1* *svm-class0*)))
    (display* "The prediction is " prediction "(" out ")"
	      ".  Correct is " (data-point-class data-point)
	      "."
	      )
    prediction))

(define (make-svm-prediction point-features)
  ;; initialize point n (training data is 0..n-1) and called the learned-func
  (let ((n (svm-get-n-points)))
    (svm-set-data-point n point-features)
    (if (not (eq? (first (svm-get-kernel-spec)) 'linear))
	(svm-set-self-dot n (dot-product n n)))
    (svm-learned-func n)
    ))

