Invent a new combination matching macro to do it all

This commit is contained in:
Stas Boukarev 2026-09-07 01:37:18 +03:00
parent 313fd89575
commit 53d0b3c4ff
3 changed files with 359 additions and 179 deletions

View file

@ -692,6 +692,252 @@
(declare (ignorable name combination args rotated))
,match-form)))))))
(defmacro combination-match2 ((node) &body clauses)
(let (bound-vars)
(labels ((invert-relation (op)
(case op
(< '>)
(> '<)
(<= '>=)
(>= '<=)))
(invertible-p (s)
(and (consp s)
(symbolp (car s))
(invert-relation (car s))
(= (length (cdr s)) 2)
(not (equal-spec (first (cdr s)) (second (cdr s))))))
(invert-spec (s)
(list (invert-relation (car s))
(second (cdr s))
(first (cdr s))))
(ensure-or (x)
(let ((specs (if (typep x '(cons (eql :or)))
(cdr x)
(list x))))
;; Expand (< x y) into (:or (< x y) (> y x))
;; And = into eq eql =
(loop for s in specs
collect s
when (invertible-p s)
collect (invert-spec s))))
(collect-spec-vars (spec)
(let (vars)
(labels ((add (s)
(when (and s
(symbolp s)
(not (keywordp s))
(not (eq s '*))
(not (eq s '&rest)))
(pushnew s vars)))
(walk (s)
(cond ((typep s '(cons (eql :type)))
(let ((var (third s)))
(add var)))
((typep s '(cons (eql :constant)))
(add (second s)))
((typep s '(cons (member :or :commutative)))
(mapc #'walk (cdr s)))
((consp s)
(mapc #'walk (cdr s)))
(t
(add s)))))
(walk spec)
(nreverse vars))))
(equal-spec (a b)
(cond ((eq a b))
((symbolp a)
(and (symbolp b)
(eq a b)))
((typep a '(cons (member :type :constant)))
(equal a b))
((and (consp a)
(consp b))
(and (equal (car a)
(car b))
(= (length a)
(length b))
(every #'equal-spec a b)))))
(expand-node (lvars specs spec body)
(let ((old-bound-vars bound-vars)
(spec (ensure-or spec)))
(labels ((gen (&optional sub)
(destructuring-bind (name . args) (pop spec)
(let* ((variable (position '&rest args))
(arg-count (or variable
(length args)))
(vars (make-gensym-list arg-count "ARG"))
(names (ensure-or name))
(commutative (and (loop for name in names
always (or (typep name '(cons (eql :commutative)))
(ir1-attributep (fun-info-attributes (fun-info-or-lose name))
commutative)))
(not (or (integerp (car (last args)))
(typep (car (last args)) '(cons (eql :constant)))
(equal-spec (first args)
(second args))))))
(names (loop for name in names
collect (if (typep name '(cons (eql :commutative)))
(second name)
name))))
(setf bound-vars old-bound-vars)
(let ((args
`(or (multiple-value-bind ,vars ,(if variable
`(check-min-args args ,arg-count)
`(check-args args ,arg-count))
(declare (ignorable ,@vars))
(when ,(car vars)
,(expand lvars specs
(lambda ()
(let ((old-bound-vars bound-vars))
(cond (commutative
(assert (= (length vars) 2))
`(or ,(expand vars args body)
,(progn
(setf bound-vars old-bound-vars)
(expand (list (second vars) (first vars)) args body))))
(t
(expand vars args body))))))))
,@(unless sub
(loop while (and spec
(subsetp (ensure-or (caar spec))
names))
collect
`(case name
,(gen t)))))))
`(,names ,args))))))
(loop while spec
collect (gen)))))
(expand (lvars specs body)
(if lvars
(let ((lvar (car lvars))
(spec (car specs)))
(flet ((match-var (name &optional allow-empty constant constant-type)
(cond ((or (eq name '*)
(and allow-empty
(not name)))
(expand (cdr lvars) (cdr specs)
body))
((member name bound-vars)
`(when ,(if constant
`(eql ,name (lvar-value ,lvar))
`(same-leaf-ref-p ,name ,lvar))
,(expand (cdr lvars) (cdr specs)
body)))
(t
(push name bound-vars)
(let ((expanded (expand (cdr lvars) (cdr specs)
body)))
`(let ((,name ,(if constant
`(lvar-value ,lvar)
lvar)))
,(if constant-type
`(when (typep ,name ',constant-type)
,expanded)
expanded)))))))
(cond ((typep spec '(cons (eql :type)))
`(when (csubtypep (lvar-type ,lvar) (specifier-type ',(second spec)))
,(match-var (third spec) t)))
((typep spec '(cons (eql :constant)))
`(when (constant-lvar-p ,lvar)
,(match-var (second spec) t t (third spec))))
((symbolp spec)
(match-var spec))
((atom spec)
`(when (lvar-value-is ,lvar ',spec)
,(expand (cdr lvars) (cdr specs)
body)))
(t
`(multiple-value-bind (name combination args) (lvar-combination/cast-name-args ,lvar)
(declare (notinline lvar-value-is))
(when combination
(case name
,@(expand-node (cdr lvars) (cdr specs) spec body))))))))
(funcall body)))
(gen-1 (clauses node)
(let ((flets nil)
(sym-forms (make-hash-table :test 'eq))
sym-order
form-groups)
(dolist (clause clauses)
(destructuring-bind (spec &body body) clause
(setf bound-vars nil)
(let* ((pattern-vars (collect-spec-vars spec))
(body-fun (gensym "MATCH-BODY"))
(matched (lambda ()
`(return-from .combination-match.
(,body-fun name combination args ,@pattern-vars))))
(branches (expand-node nil nil spec matched)))
(push `(,body-fun (name combination args ,@pattern-vars)
(declare (ignorable name combination args ,@pattern-vars))
(let ((new (progn ,@body)))
(when new
(combination-match-transform .node. ',pattern-vars new ,@pattern-vars))))
flets)
(dolist (branch branches)
(destructuring-bind (names form) branch
(dolist (name names)
(unless (gethash name sym-forms)
(push name sym-order))
(push form (gethash name sym-forms))))))))
(setf sym-order (nreverse sym-order))
(dolist (sym sym-order)
(let* ((forms (nreverse (gethash sym sym-forms)))
(entry (assoc forms form-groups :test #'equal)))
(if entry
(push sym (cdr entry))
(push (cons forms (list sym)) form-groups))))
(let ((case-branches
(loop for (forms . syms) in (nreverse form-groups)
collect `(,(nreverse syms)
,(if (cdr forms)
`(or ,@forms)
(car forms))))))
`(let ((.node. ,node))
(flet ,flets
(multiple-value-bind (name combination args) (combination/cast-name-args .node.)
(declare (ignorable name combination args))
(case name
,@case-branches))))))))
(let ((dest (member :dest clauses)))
`(progn
(block .combination-match.
,(gen-1 (ldiff clauses dest) node)
,(when dest
(gen-1 (cdr dest) `(node-dest ,node))))
;; Matching multiple combinations may not get reoptimized,
;; so register to get another chance
(delay-ir1-transform node :ir1-phases)
(give-up-ir1-transform))))))
(defun combination-match-transform (combination vars form &rest lvars)
(loop for var in vars
for lvar in lvars
when (lvar-p lvar) ;; ignore constants
collect lvar into lvars*
and
collect var into vars*
finally (setf vars vars*
lvars lvars*))
(let ((old-args (combination-args combination)))
(loop for lvar in lvars
do
(extract-lvar lvar combination))
(loop for arg in old-args
unless (member arg lvars :test #'eq)
do (flush-dest arg))
(setf (combination-args combination)
lvars)
(when *show-transforms-p*
(show-transform :combination-match 'x form combination))
(transform-call combination
`(lambda ,vars
(declare (ignorable ,@vars))
,(unless (eq form :nil)
form))
'combination-match2))
(throw 'give-up-ir1-transform :none))
(defun erase-node-type (node type &optional nth-value erase-calls)
(setf (node-derived-type node)
(cond ((eq type t)

View file

@ -3199,39 +3199,28 @@
(give-up-ir1-transform))))
(deftransform integer-length ((x) (integer) * :node node :important nil)
(block nil
(or (combination-case x
(lognot (*)
(splice-fun-args x :any 1)
(give-up-ir1-transform)))
(let ((dest (sole-node-dest node)))
(combination-case (nil :node dest)
(< (* (constant (integer 0 4096)))
(let ((k (lvar-value (second args))))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
(cond ((zerop k)
`(progn nil))
(combination-match2 (node)
((integer-length (lognot x))
`(integer-length x))
:dest
((< (integer-length x) (:constant c (integer 0 4096)))
(cond ((zerop c)
:nil)
((lvar-subtypep x unsigned-byte)
`(< x ,(ash 1 (1- k))))
`(< x ,(ash 1 (1- c))))
(t
`(typep x '(signed-byte ,k))))))
(> (* (constant (integer 0 4096)))
(let ((k (lvar-value (second args))))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
`(typep x '(signed-byte ,c)))))
((> (integer-length x) (:constant c (integer 0 4096)))
(cond ((lvar-subtypep x unsigned-byte)
`(not (< x ,(ash 1 k))))
((zerop k)
`(not (< x ,(ash 1 c))))
((zerop c)
(cond ((word-sized-lvar-p x)
`(>= (logand most-positive-word (+ x 1)) 2))
(t
'(not (or (eq x 0) (eq x -1))))))
(t
`(not (typep x '(signed-byte ,(1+ k))))))))
(eq (* (constant (integer 0 4096)))
(let* ((c (lvar-value (second args)))
(transform
`(not (typep x '(signed-byte ,(1+ c)))))))
((eq (integer-length x) (:constant c (integer 0 4096)))
(cond
((zerop c)
(cond ((lvar-subtypep x unsigned-byte)
@ -3249,7 +3238,7 @@
((<= c sb-vm:n-word-bits)
`(= (ash x ,(- 1 c)) 1))
(t
'(progn nil))))
:nil)))
((lvar-subtypep x signed-word)
(cond ((= c (1- sb-vm:n-word-bits))
`(not (typep x '(signed-byte ,(1- sb-vm:n-word-bits)))))
@ -3257,13 +3246,7 @@
`(let ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits)))))
(= (ash x ,(- 1 c)) 1)))
(t
'(progn nil)))))))
(when transform
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
transform)))))
(delay-ir1-transform node :ir1-phases)
(give-up-ir1-transform))))
:nil)))))))
(defoptimizer (%bignum-length derive-type) ((x))
(one-arg-derive-type
@ -3332,126 +3315,74 @@
(specifier-type '(integer 1)))))))))
(deftransform logcount ((x) (integer) * :node node :important nil)
(let ((dest (sole-node-dest node)))
(block nil
(or (combination-case x
(lognot (*)
(splice-fun-args x :any 1)
(give-up-ir1-transform))
(ash ((type unsigned-byte) (type unsigned-byte))
(splice-fun-args x :any #'first)
(give-up-ir1-transform))
(* ((type unsigned-byte) (constant unsigned-byte))
(when (= (logcount (lvar-value (second args))) 1)
(splice-fun-args x :any #'first)
(give-up-ir1-transform)))
(* ((type (integer * 0)) (constant (integer * -1)))
(when (= (logcount (abs (lvar-value (second args)))) 1)
(splice-fun-args x :any #'first)
`(logcount (%negate x)))))
(combination-match (:sole-dest node)
(:or (> (logcount x) (integer-length x))
(< (integer-length x) (logcount x)))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity (if (eq name 'integer-length) 0 1) dest 2)
(return nil))
(combination-match (:sole-dest node)
(:or (< (logcount x) (integer-length x))
(> (integer-length x) (logcount x)))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity (if (eq name 'integer-length) 0 1) dest 2)
(combination-match2 (node)
((logcount (lognot x))
`(logcount x))
((logcount (ash (:type unsigned-byte x) (:type unsigned-byte *)))
`(logcount x))
((logcount (* (:type unsigned-byte x) (:constant c unsigned-byte)))
(when (= (logcount c) 1)
`(logcount x)))
((logcount (* (:type (integer * 0) x) (:constant c (integer * -1))))
(when (= (logcount (abs c)) 1)
`(logcount (%negate x))))
:dest
((< (logcount x) (integer-length x))
`(not (eq (logcount x) (integer-length x))))
(unless (word-sized-lvar-p x)
;; (= (logcount signed) 0) => (or (eq x 0) (eq x -1))
(combination-case (nil :node dest)
(eq (* 0)
((> (logcount x) (integer-length x))
:nil)
((eq (logcount x) 0)
(unless (word-sized-lvar-p x)
(delay-ir1-transform node :ir1-phases)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
'(or (eq x 0) (eq x -1))))))
(give-up-ir1-transform)))))
(deftransform logcount ((x) (unsigned-byte) * :node node :important nil)
(or (let ((dest (sole-node-dest node)))
(combination-case (nil :node dest)
;; (= (logcount unsigned) 0) => (= x 0)
(eq (* 0)
(erase-node-type node (lvar-single-value-type x))
'x)))
(give-up-ir1-transform)))
(combination-match2 (node)
:dest
((eq (logcount x) 0)
x)))
(deftransform logcount ((x) (signed-word) * :node node :important nil)
(or
(let ((dest (sole-node-dest node)))
(cond
((or
(combination-match (:sole-dest node)
(eq (logcount x) (integer-length x))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity (if rotated 1 0) dest 2)
(delay-ir1-transform node :ir1-phases) ;; unsigned transforms are better
(combination-match2 (node)
:dest
((eq (logcount x) (integer-length x))
`(let ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits)))))
(not (logtest x (+ x 1)))))))
(t
(delay-ir1-transform node :ir1-phases)
(combination-case (nil :node dest)
(eq (* 0)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
(not (logtest x (+ x 1)))))
((eq (logcount x) 0)
`(< (logand most-positive-word (+ x 1)) 2))
(eq (* 1)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
((eq (logcount x) 1)
(if (lvar-intersectp x (integer -1 0))
`(let* ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits))))
(1-x (- x 1)))
(> (logxor x 1-x) 1-x))
`(let ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits)))))
(not (logtest x (- x 1))))))
(< (* 2)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
((< (logcount x) 2)
`(let ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits)))))
(not (logtest x (- x 1)))))
(> (* 1)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
((> (logcount x) 1)
`(let ((x (logxor x (ash x ,(- 1 sb-vm:n-word-bits)))))
(logtest x (- x 1))))))))
(give-up-ir1-transform)))
(logtest x (- x 1))))))
(deftransform logcount ((x) (word) * :node node :important nil)
(or (let ((dest (sole-node-dest node)))
(cond
(combination-match2 (node)
:dest
(((:or eq eql =) (logcount x) (integer-length x))
;; Don't delay, integer-length is transformed to clz/cls on arm64
((combination-match (:sole-dest node)
(eq (logcount x) (integer-length x))
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity (if rotated 1 0) dest 2)
`(not (logtest x (+ x 1)))))
(t
(delay-ir1-transform node :ir1-phases)
(combination-case (nil :node dest)
;; (= (logcount unsigned) 0) => (= x 0)
(eq (* 0)
(erase-node-type node (lvar-single-value-type x))
'x)
(eq (* 1)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
`(not (logtest x (+ x 1))))
((eq (logcount x) 0)
`(eq x 0))
((eq (logcount x) 1)
(if (lvar-intersectp x (eql 0))
`(let ((1-x (logand most-positive-word (- x 1))))
(> (logxor x 1-x) 1-x))
`(not (logtest x (- x 1)))))
(< (* 2)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
((< (logcount x) 2)
`(not (logtest x (- x 1))))
(> (* 1)
(erase-node-type node (values-specifier-type '(values boolean &optional)))
(transform-to-identity 0 dest 2)
`(logtest x (- x 1)))))))
(give-up-ir1-transform)))
((> (logcount x) 1)
`(logtest x (- x 1)))))
(defoptimizer (isqrt derive-type) ((x))
(one-arg-derive-type

View file

@ -1962,5 +1962,8 @@
(#(6A3E03D5 73D42188 937BB764 A4528420 D0F360C2 D1F36255 D5F368A1 D7F36BC7)
"(ABS * ASH TRUNCATE / %NEGATE + -)"
"((& (- val (>> val 28)) 7))")
(#(56B428A2 C4F34DDE C6F35104 C7F35297 DC33E6FC)
"(> < = EQL EQ)"
"((& (^ (>> val 2) (>> val 7)) 7))")
)
;; EOF