
(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)

;;; This is a simple record definition capability that is portable (or
;;; as portable as eval, which is not much).
(define (defstruct name var components)
  (define (symbol-append . l)
    (string->symbol (apply string-append (map symbol->string l))))
  (define (index-list l i)		; list of ints starting with i
    (if (null? l) '()
	(cons i (index-list (cdr l) (+ i 1)))))
  ;; Create the global variable to hold the data
  (scheme-eval `(define ,var #f))
  ;; This is the make-x function, which initializes the global var
  (scheme-eval `(define (,(symbol-append 'make '- name)) 
	   (set! ,var (make-vector ,(+ 1 (length components))))
	   (vector-set! ,var 0 ',name)))
  ;; Create each of the access and set functions
  (for-each (lambda (c i)		; c if component name
	      ;; Create name-c
	      (scheme-eval `(define (,(symbol-append name '- c)) 
		       (vector-ref ,var ,i)))
	      ;; Create name-set-c
	      (scheme-eval `(define (,(symbol-append name '- 'set '- c) x) 
		       (vector-set! ,var ,i x))))
	    components
	    (index-list components 1)))

;;; Constructs the svm-x and svm-set-x functions for each of the
;;; components of the SVM data structure.  The disadvantage is that
;;; these generated functions will not get compiled.
(defstruct 'svm '*svm* 
  '(n-points n-supports dim kernel-spec alphas weights threshold points 
	     targets class0 class1))

;;; Used during training to avoid repeated work.
(define *svm-error-cache* #f)
(define *svm-precomputed-self-dot-product* #f)

;;; Some nicer ways of getting at the components of vector parameters
(define (svm-get-target i) (vector-ref (svm-targets) i))
(define (svm-set-target i v) (vector-set! (svm-targets) i v))

(define (svm-get-alpha i) (vector-ref (svm-alphas) i))
(define (svm-set-alpha i a) (vector-set! (svm-alphas) i a))

(define (svm-get-point i) (vector-ref (svm-points) i))
(define (svm-set-point i p) (vector-set! (svm-points) i p))

(define (svm-get-weight i) (vector-ref (svm-weights) i))
(define (svm-set-weight i w) (vector-set! (svm-weights) i w))

(define (svm-get-error-cache i) (vector-ref *svm-error-cache* i))
(define (svm-set-error-cache i e) (vector-set! *svm-error-cache* i e))

(define (svm-get-self-dot i) (vector-ref *svm-precomputed-self-dot-product* i))
(define (svm-set-self-dot i d) (vector-set! *svm-precomputed-self-dot-product* i d))

(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-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-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-n-points)))
		 (if *svm-verbose* (display* 'heuristc3 " " k " " end))
		 1)
		(else
		 (heuristic3 (+ k 1) end)))))

    (if *svm-verbose*
	(display* "alpha=" (svm-alphas) "b=" (svm-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-n-points))
		      ))
	      (heuristic2 k0 (+ (svm-n-points) k0)))
	    (let ((k0 (random (svm-n-points))
		      ))
	      (heuristic3 k0 (+ (svm-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-kernel-spec)) 'linear)
      (do ((i 0 (1+ i))
	   (p1 (svm-get-point i1) (cdr p1))
	   (p2 (svm-get-point i2) (cdr p2)))
	  ((= i (svm-dim)))
	(svm-set-weight i 
			(+ (svm-get-weight i)
			   (+ (* t1 (first p1)) (* t2 (first p2)))))))
  (do ((i 0 (1+ i)))
      ((= i (svm-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-threshold))
	(bnew
	 (if (and (> a1 0) (< a1 *svm-max-alpha*))
	     (+ (svm-threshold) e1 (* t1 k11) (* t2 k12))
	     (if (and (> a2 0) (< a2 *svm-max-alpha*))
		 (+ (svm-threshold) e2 (* t1 k12) (* t2 k22))
		 (let ((b1 (+ (svm-threshold) e1 (* t1 k11) (* t2 k12)))
		       (b2 (+ (svm-threshold) e2 (* t1 k12) (* t2 k22))))
		   (/ (+ b1 b2) 2))))))
    (svm-set-threshold bnew)
    ;; return delta-b
    (- bnew bold)))

(define (learned-func-linear k) 
  (do ((i 0 (1+ i))
       (pk (svm-get-point k) (cdr pk))
       (s 0 (+ s (* (svm-get-weight i) (first pk)))))
      ((= i (svm-dim)) (- s (svm-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-n-supports)) (- s (svm-threshold))))
  )

(define (svm-learned-func k)
  (let ((kernel-spec (svm-kernel-spec)))
    (cond ((eq? (first kernel-spec) 'linear) (learned-func-linear k))
	  ((memq (first kernel-spec) '(gauss)) (learned-func-nonlinear k))
	  (else (error "Unsupported kernel: " kernel-spec)))))

(define (dot-product i1 i2)
  (apply + (map * (svm-get-point i1) (svm-get-point i2))))

(define (rbf-kernel i1 i2)
  ;; (p1 - p2).(p1 - p2) = p1.p1 -2 p1.p2 + p2.p2
  (exp (/ (- (+ (svm-get-self-dot i1)
		(* -2 (dot-product i1 i2))
		(svm-get-self-dot i2)))
	  ;; kernel spec is (rbf two-sigm-squared)
	  (second (svm-kernel-spec)))))

(define (svm-kernel-func i1 i2)
  (let ((kernel-spec (svm-kernel-spec)))
    (cond ((eq? (first kernel-spec) 'linear) (dot-product i1 i2))
	  ((eq? (first kernel-spec) 'gauss) (rbf-kernel i1 i2))
	  (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)))
		  (svm-set-class0 0)
		  (svm-set-class1 1))
		 ((member classes '((-1 1) (-1.0 1.0) (-1 0) (1.0 -1.0)))
		  (svm-set-class0 -1)
		  (svm-set-class1 1))
		 (else
		  (svm-set-class0 (first classes))
		  (svm-set-class1 (second classes)))))
	  (else 
	   (error "Only two classes allowed, but we have" classes)))))

(define (svm-setup data kernel-spec)
  ;; Set up the SVM data structures
  (make-svm)
  (let ((n (length data)))
    (svm-set-n-points n)
    (svm-set-n-supports n))
  (svm-set-dim (length (data-point-features (first data))))
  (svm-set-kernel-spec kernel-spec)
  (svm-set-alphas (make-vector (svm-n-points) 0.0))
  (if (eq? (first kernel-spec) 'linear)
      (svm-set-weights (make-vector (svm-dim) 0.0)))
  (svm-set-threshold 0)
  (svm-set-points (make-vector (svm-n-points)))
  (svm-set-targets (make-vector (svm-n-points)))
  (svm-set-classes data)
  (do ((i 0 (+ i 1))
       (p data (cdr p)))
      ((null? p))
    (svm-set-point i (data-point-features (car p)))
    (let ((class (data-point-class (car p))))
      (svm-set-target i (if (eq? class (svm-class0)) -1 1)))
    )
  (cond ((not (eq? (first kernel-spec) 'linear))
	 (set! *svm-precomputed-self-dot-product* (make-vector (svm-n-points)))
	 (do ((i 0 (+ i 1)))
	     ((= i (svm-n-points)))
	   (svm-set-self-dot i (dot-product i i)))))
  )

(define (svm-train-setup data kernel-func)
  (svm-setup data kernel-func)
  (set! *svm-error-cache* (make-vector (svm-n-points) 0.)))

(define (svm-train data kernel-spec)
  (svm-train-setup data kernel-spec)
  (do ((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-n-points)))
	  (set! num-changed (+ num-changed (examine-example k))))
	(do ((k 0 (+ 1 k)))
	    ((= k (svm-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-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)))))
      )
    )
  
  (do ((i 0 (+ i 1))
       (errors 0))
      ((= i (svm-n-points))
       (display* "Got " errors " prediction errors.  Error rate = " 
		 (exact->inexact (/ errors (svm-n-points)))))
    (if (< (* (svm-learned-func i) (svm-get-target i)) 0) 
	(set! errors (+ 1 errors))))
  
  'done
  )

